72 Commits
Author SHA1 Message Date
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
148 changed files with 10583 additions and 5163 deletions
+3 -1
View File
@@ -4,6 +4,8 @@
# Allow necessary files
!astrai/
!scripts/
!assets/
!docs/
!csrc/
!setup.py
!pyproject.toml
!README.md
+18 -10
View File
@@ -26,22 +26,30 @@ jobs:
if-no-files-found: error
build-cuda-linux:
name: Build CUDA wheel (Linux)
name: Build CUDA wheel (Linux, ${{ matrix.cuda_tag }})
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:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: "3.12"
- name: Install torch (CUDA 12.8)
- name: Install torch (${{ matrix.cuda_tag }})
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
with:
cuda: "12.8.0"
cuda: "${{ matrix.cuda_ver }}"
- name: Build wheel (with CUDA kernels)
run: |
@@ -49,7 +57,7 @@ jobs:
- uses: actions/upload-artifact@v4
with:
name: cuda-wheel-linux
name: cuda-wheel-linux-${{ matrix.cuda_tag }}
path: dist/*.whl
if-no-files-found: error
@@ -66,10 +74,11 @@ jobs:
name: pure-wheel
path: release-assets/pure
- name: Download CUDA wheel
- name: Download CUDA wheels (all variants)
uses: actions/download-artifact@v4
with:
name: cuda-wheel-linux
pattern: cuda-wheel-linux-*
merge-multiple: true
path: release-assets/cuda
- name: Verify release assets
@@ -79,8 +88,7 @@ jobs:
pure_wheels=(release-assets/pure/*.whl)
cuda_wheels=(release-assets/cuda/*.whl)
test "${#pure_wheels[@]}" -eq 1
test "${#cuda_wheels[@]}" -eq 1
test "$(basename "${pure_wheels[0]}")" != "$(basename "${cuda_wheels[0]}")"
test "${#cuda_wheels[@]}" -ge 1
- name: Create release & upload assets
uses: softprops/action-gh-release@v2
+1 -1
View File
@@ -24,7 +24,7 @@
!/.dockerignore
!/Dockerfile
!/docker-compose.yml
!/assets/**
!/docs/**
!/CONTRIBUTING.md
!/LICENSE
!/pyproject.toml
+9 -7
View File
@@ -20,9 +20,6 @@ Run the following checks **in order** — CI will reject if any fail.
ruff format .
```
> **Note**: `ruff format` may rename parameters (e.g. `mask` → `attn_mask`).
> Always review the diff after formatting.
### 2. Import sorting
```bash
@@ -44,7 +41,7 @@ python -u -m pytest tests/ -v
> 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:
@@ -52,12 +49,17 @@ If you have Git Bash available:
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
```
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)
```
@@ -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 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 |
## Submitting Changes
+12 -2
View File
@@ -1,8 +1,16 @@
# 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
FROM ubuntu:24.04 AS builder
ARG CUDA_TAG=cu128
WORKDIR /app
# 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 astrai/ ./astrai/
COPY csrc/ ./csrc/
COPY setup.py .
COPY pyproject.toml .
RUN pip install --no-cache-dir --upgrade pip \
&& 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
FROM ubuntu:24.04 AS production
@@ -43,7 +53,7 @@ ENV PATH="/opt/venv/bin:$PATH"
# Copy application code
COPY astrai/ ./astrai/
COPY scripts/ ./scripts/
COPY assets/ ./assets/
COPY docs/ ./docs/
COPY pyproject.toml .
COPY README.md .
+19 -12
View File
@@ -1,6 +1,6 @@
<div align="center">
<img src="assets/images/logo.png" width="auto" alt="Logo">
<img src="docs/images/logo.png" width="auto" alt="Logo">
<p>
<strong>A lightweight Transformer training & inference framework</strong>
</p>
@@ -17,7 +17,7 @@
<div align="center">
<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/discussions">Discussions</a> •
<a href="https://huggingface.co/ViperEkura">HuggingFace</a>
@@ -56,6 +56,8 @@ End-to-end walkthrough in 5 steps:
**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
git clone https://github.com/ViperEkura/AstrAI.git
cd AstrAI
@@ -132,7 +134,7 @@ Check out the demos in the `scripts/demo/` folder:
# Download model weights (required before running demos)
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
# Type your message after >>, type !exit to quit
@@ -183,7 +185,7 @@ docker run --gpus all -v /path/to/data:/data -it astrai:latest
# Docker Compose (GPU, default)
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
```
@@ -213,18 +215,23 @@ curl -X POST http://localhost:8000/v1/messages \
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
| Document | Description |
|----------|-------------|
| [CLI Reference](./assets/docs/params.md) | Parameters for all CLI tools (train, server, generate, preprocess) |
| [Architecture](./assets/docs/architecture.md) | System architecture, class diagram & design patterns |
| [Training](./assets/docs/training.md) | Training loop, strategies & formulas |
| [Inference](./assets/docs/inference.md) | KVCache, continuous batching, sampling & HTTP API |
| [Data Flow](./assets/docs/dataflow.md) | Data pipeline, storage backends & dataset architecture |
| [Preprocessing](./assets/docs/preprocessing.md) | Declarative JSON-driven data preprocessing |
| [Get Started](./docs/get-started.md) | Installation and quickstart |
| [CLI Reference](./docs/guides/params.md) | Parameters for all CLI tools (train, server, generate, preprocess) |
| [Preprocessing](./docs/guides/preprocessing.md) | Declarative JSON-driven data preprocessing |
| [Training](./docs/guides/training.md) | Training loop, strategies & formulas |
| [Inference](./docs/guides/inference.md) | KVCache, continuous batching, sampling & HTTP API |
| [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
@@ -251,4 +258,4 @@ This project is licensed under the [GPL-3.0 License](LICENSE).
<div align="center">
<em>A lightweight Transformer framework designed for both high performance and ease of use.</em>
</div>
</div>
-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_position_embeddings=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
+29 -1
View File
@@ -1,6 +1,9 @@
__version__ = "1.3.11"
__version__ = "1.3.12"
__author__ = "ViperEkura"
import logging
import os
from astrai.config import (
AutoRegressiveLMConfig,
BaseModelConfig,
@@ -53,6 +56,30 @@ from astrai.trainer import (
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__ = [
"AutoRegressiveLM",
"AutoRegressiveLMConfig",
@@ -94,5 +121,6 @@ __all__ = [
"only_on_rank",
"run_server",
"sample",
"setup_logging",
"spawn_parallel_fn",
]
+20 -80
View File
@@ -1,92 +1,32 @@
import json
from dataclasses import MISSING, dataclass, fields
from dataclasses import asdict
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:
def to_dict(self) -> Dict[str, Any]:
d = {}
for fld in fields(self):
v = getattr(self, fld.name)
if isinstance(v, (str, int, float, bool)):
d[fld.name] = v
elif v is None:
d[fld.name] = None
elif isinstance(v, (dict, list, tuple)):
try:
val = list(v) if isinstance(v, tuple) else v
json.dumps(val)
d[fld.name] = val
except (TypeError, ValueError):
pass
elif isinstance(v, BaseConfig):
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
result = {}
for k, v in asdict(self).items():
if isinstance(v, tuple):
v = list(v)
try:
json.dumps(v)
result[k] = v
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
return result
@classmethod
def from_dict(cls, d: Dict[str, Any]) -> Self:
hints = get_type_hints(cls)
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
return cls(**d)
@classmethod
def from_file(cls, path: Union[str, Path]) -> Self:
+106 -11
View File
@@ -1,9 +1,14 @@
from dataclasses import dataclass
from typing import Any, Dict, Optional
from pydantic import field_validator
from pydantic.dataclasses import dataclass
from astrai.config.base import BaseConfig
from astrai.factory import BaseFactory
_ATTN_TYPES = frozenset({"gqa", "mla"})
_FFN_TYPES = frozenset({"mlp", "moe"})
class ConfigFactory(BaseFactory[BaseConfig]):
"""Factory that dispatches config classes by ``model_type``."""
@@ -17,7 +22,12 @@ class ConfigFactory(BaseFactory[BaseConfig]):
@dataclass
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
neftune_alpha: float = 0.0
@@ -26,7 +36,39 @@ class BaseModelConfig(BaseConfig):
@dataclass
@ConfigFactory.register("autoregressive_lm")
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
hidden_size: Optional[int] = None
@@ -34,49 +76,102 @@ class AutoRegressiveLMConfig(BaseModelConfig):
rms_norm_eps: Optional[float] = None
intermediate_size: Optional[int] = None
tie_word_embeddings: Optional[bool] = None
max_position_embeddings: Optional[int] = None
rope_theta: Optional[float] = None
rope_scaling: Optional[dict] = None
attn_type: str = "gqa"
num_attention_heads: Optional[int] = None
num_key_value_heads: Optional[int] = None
use_qk_norm: Optional[bool] = None
use_gated_attention: Optional[bool] = None
kv_lora_rank: Optional[int] = None
qk_nope_head_dim: Optional[int] = None
qk_rope_head_dim: Optional[int] = None
ffn_type: str = "mlp"
n_routed_experts: Optional[int] = None
n_shared_experts: Optional[int] = None
n_activated_experts: Optional[int] = 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
@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
@ConfigFactory.register("embedding")
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
hidden_size: Optional[int] = None
num_hidden_layers: Optional[int] = None
rms_norm_eps: Optional[float] = None
intermediate_size: Optional[int] = None
max_position_embeddings: Optional[int] = None
rope_theta: Optional[float] = None
rope_scaling: Optional[dict] = None
attn_type: str = "gqa"
num_attention_heads: Optional[int] = None
num_key_value_heads: Optional[int] = None
use_qk_norm: Optional[bool] = None
use_gated_attention: Optional[bool] = None
ffn_type: str = "mlp"
pooling_type: Optional[str] = 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
+85 -45
View File
@@ -5,11 +5,19 @@ modes, both driven declaratively through ``input.sections`` or
``input.sources``.
"""
from dataclasses import dataclass, field
from dataclasses import field
from typing import Dict, List, Optional
from pydantic import field_validator
from pydantic.dataclasses import dataclass
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
class InputConfig(BaseConfig):
@@ -25,6 +33,10 @@ class InputConfig(BaseConfig):
"chosen": {"sections": [{"field": "chosen", ...}]},
"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
@@ -33,34 +45,17 @@ class InputConfig(BaseConfig):
@dataclass
class ProcessingConfig(BaseConfig):
"""Processing configuration.
"""Processing configuration for tokenization and packing.
Parameters
----------
max_seq_len : int
Maximum sequence length (default: 2048).
min_chars : int
Minimum number of characters to keep (default: 50).
max_chars : int
Maximum number of characters to keep (default: 2_000_000).
max_items : Optional[int]
Maximum number of items to process (default: None, unlimited).
batch_size : int
Number of records tokenized together (default: 256).
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.
Args:
max_seq_len (int): Maximum sequence length. Defaults to 2048.
min_chars (int): Minimum number of characters to keep. Defaults to 50.
max_chars (int): Maximum number of characters to keep. Defaults to 2_000_000.
max_items (Optional[int]): Maximum number of items to process, None=unlimited. Defaults to None.
batch_size (int): Number of records tokenized together. Defaults to 256.
packing_strategy (str): How to pack sequences: 'simple', 'bfd', or 'bfd_split'. Defaults to "simple".
max_packed_len (int): Maximum length of a packed bin. Defaults to 8192.
truncation_mode (str): How to truncate over-length sequences: 'keep_start' or 'keep_end'. Defaults to "keep_start".
"""
max_seq_len: int = 2048
@@ -72,27 +67,45 @@ class ProcessingConfig(BaseConfig):
max_packed_len: int = 8192
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
class OutputConfig(BaseConfig):
"""Output configuration.
"""Output configuration for storage.
Parameters
----------
domain_key : Optional[str]
Domain key for the output store (default: None).
storage_format : str
Storage format, one of ``"bin"``, ``"jsonl"`` (default: ``"bin"``).
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).
Args:
domain_key (Optional[str]): Domain key for the output store. Defaults to None.
storage_format (str): Storage format: 'bin' or 'jsonl'. Defaults to "bin".
max_tokens_per_shard (int): Maximum tokens per shard before splitting. Defaults to 100_000_000.
dtype (Dict[str, str]): Per-key dtype overrides, e.g. {"input_ids": "int32"}. Defaults to {}.
position_ids_mode (str): Position ids mode: 'none', 'doc_reset', or 'continuous'. Defaults to "doc_reset".
"""
domain_key: Optional[str] = None
@@ -101,9 +114,36 @@ class OutputConfig(BaseConfig):
dtype: Dict[str, str] = field(default_factory=dict)
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
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
input: InputConfig = field(default_factory=InputConfig)
mask: Dict[str, str] = field(default_factory=dict)
+195 -158
View File
@@ -1,7 +1,9 @@
from dataclasses import dataclass, field, fields
from dataclasses import field
from typing import Any, Callable, Dict, List, Optional
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.lr_scheduler import LRScheduler
from torch.utils.data import Dataset
@@ -9,173 +11,208 @@ from torch.utils.data import Dataset
from astrai.config.base import BaseConfig
from astrai.model.components.lora import LoRAConfig
def required(**kw):
return {"required": True, **kw}
_TRAIN_TYPES = frozenset({"seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"})
_PARALLEL_MODES = frozenset({"none", "ddp", "fsdp"})
_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):
# basic setting
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=1.0,
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."},
)
"""Training configuration.
# checkpoint setting
start_epoch: int = field(default=0, metadata={"help": "Start epoch for training."})
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."},
)
Combines hyperparameters with runtime objects (model_fn, dataset, etc.).
Only JSON-serializable fields are written to checkpoint meta via to_dict().
# lora setting
lora: Optional[LoRAConfig] = field(
default=None,
metadata={"help": "LoRA config. None means full fine-tuning."},
)
Args:
model_fn (Callable[[], nn.Module]): Model factory for training.
strategy (str): Training strategy (seq, sft, dpo, grpo, online_*).
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
log_dir: str = field(
default="./checkpoint/logs", metadata={"help": "Directory for metric logs."}
)
metrics: List[str] = field(
default_factory=lambda: ["loss", "lr", "grad_norm"],
metadata={"help": "Metrics to record during training."},
)
model_fn: Callable[[], nn.Module]
strategy: str
dataset: Dataset
optimizer_fn: Callable[[nn.Module], Optimizer]
scheduler_fn: Callable[[Optimizer], LRScheduler]
optimizer_name: Optional[str] = None
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
random_seed: int = field(default=3407, metadata={"help": "Random seed."})
num_workers: int = field(
default=0, metadata={"help": "Number of workers for dataloader."}
)
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)."},
)
start_epoch: int = 0
start_samples: int = 0
ckpt_dir: str = "./checkpoint"
ckpt_interval: int = 5000
# distributed training
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)."},
)
lora: Optional[LoRAConfig] = None
# others
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)."},
)
metrics: List[str] = field(default_factory=lambda: ["loss", "lr", "grad_norm"])
# online rollout
rollout_interval: int = field(
default=512,
metadata={"help": "Number of optimizer steps between online rollouts."},
)
rollout_temperature: float = field(
default=0.7, metadata={"help": "Sampling temperature for online rollout."}
)
rollout_top_k: int = field(
default=0, metadata={"help": "Top-k filtering for online rollout (0=disable)."}
)
rollout_top_p: float = field(
default=0.9,
metadata={"help": "Top-p (nucleus) filtering for online rollout."},
)
rollout_max_tokens: int = field(
default=1024,
metadata={"help": "Maximum generated tokens per response in rollout."},
)
reward_model_fn: Optional[Callable] = field(
default=None,
metadata={
"help": "Factory for reward model (required for online RL strategies)."
},
)
random_seed: int = 3407
num_workers: int = 0
prefetch_factor: Optional[int] = None
pin_memory: bool = False
collate_fn: Optional[Callable[[List[Any]], Any]] = None
executor_kwargs: Dict[str, Any] = field(
default_factory=dict,
metadata={"help": "Extra kwargs passed to ExecutorFactory.create()."},
)
extra_kwargs: Dict[str, Any] = field(
default_factory=dict, metadata={"help": "Other arguments."}
)
nprocs: int = 1
backend: str = "nccl"
master_addr: str = "localhost"
master_port: str = "29500"
parallel_mode: str = "none"
start_method: str = "spawn"
def __post_init__(self):
self.validate()
device_type: str = "cuda"
val_dataset: Optional[Dataset] = None
val_split: Optional[float] = None
val_step: int = 1000
neftune_alpha: float = 0.0
moe_aux_loss_coef: float = 0.01
def validate(self):
for fld in fields(self):
if fld.metadata.get("required") and getattr(self, fld.name) is None:
raise ValueError(f"TrainConfig.{fld.name} is required but got None.")
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.storage import (
H5Store,
JsonlStore,
MmapStore,
Recordable,
@@ -17,9 +16,7 @@ from astrai.dataset.storage import (
)
from astrai.serialization import (
load_bin,
load_h5,
save_bin,
save_h5,
)
__all__ = [
@@ -31,12 +28,9 @@ __all__ = [
"Streamable",
"Recordable",
"StoreFactory",
"H5Store",
"MmapStore",
"JsonlStore",
"detect_format",
"save_h5",
"load_h5",
"save_bin",
"load_bin",
"RDSampler",
+3 -3
View File
@@ -314,7 +314,7 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
stream datasets (SEQ/SFT). Record datasets ignore it.
stride: Stride between consecutive stream samples
(default: same as *window_size*).
storage_type: Storage backend ("h5", "bin", "jsonl") or
storage_type: Storage backend ("bin", "jsonl") or
None for auto-detection.
tokenizer_path: Path to tokenizer for lazy JSONL
tokenisation (record datasets only).
@@ -384,7 +384,7 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
"""Build an on-the-fly tokenisation processor if applicable.
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.
"""
if tokenizer_path is None or storage_type != "jsonl":
@@ -451,7 +451,7 @@ class DPODataset(BaseDataset):
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.
- **Raw JSONL** (``tokenizer_path=...``): builds a lazy processor
via :func:`dpo_processor` that tokenises on the fly — no packing,
+7 -45
View File
@@ -10,7 +10,6 @@ Architecture (composition over inheritance):
Streamable (mixin) — raw token slice fetch(begin, end, keys)
Recordable (mixin) — raw record slice fetch_record(idx, keys)
H5Store(Store, Streamable, Recordable)
MmapStore(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).
``segments_are_records`` (class attribute on each Store subclass)
tells ``_normalize`` whether segments are inherently per-record (H5/
JSONL) or opaque shards (bin). Record access for bin relies on
``_offsets`` instead.
tells ``_normalize`` whether segments are inherently per-record (JSONL)
or opaque shards (bin). Record access for bin relies on ``_offsets``
instead.
:class:`JsonlStore` supports a lazy mode (``processor=fn``) that keeps
raw records and defers tokenisation to ``fetch_record`` — used by DPO
@@ -62,7 +61,6 @@ from astrai.preprocessing.transform import TokenizeTransform
from astrai.serialization import (
load_bin,
load_bin_offsets,
load_h5,
)
logger = logging.getLogger(__name__)
@@ -83,19 +81,10 @@ def detect_format(load_path: str) -> str:
root = Path(load_path)
if root.is_file():
suffix = root.suffix.lower()
if suffix in (".h5", ".hdf5"):
return "h5"
if suffix == ".jsonl":
return "jsonl"
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)]
if bin_files:
has_meta = (root / "meta.json").exists() or len(
@@ -185,7 +174,7 @@ class Store(ABC):
"""Number of records available via :meth:`fetch_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
@@ -269,7 +258,7 @@ class Store(ABC):
Record mode: if *offsets* is provided (bin layout),
``_offsets[key]`` stores cumulative per-record offsets into the
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.
Nested keys (GRPO ``responses``/``masks`` as
@@ -305,7 +294,7 @@ class Store(ABC):
logger.warning(
"Key '%s' has %d segments with offsets — record mode "
"disabled for this key (multi-shard bin+offsets not "
"supported). Merge shards or use H5/JSONL.",
"supported). Merge shards or use JSONL.",
key,
len(segs),
)
@@ -330,7 +319,7 @@ class Streamable:
Stateless trait relying on ``self._data``, ``self._cum``,
``self._length`` maintained by :class:`Store`. Stream mode is
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.
"""
@@ -415,33 +404,6 @@ class StoreFactory(BaseFactory["Store"]):
"""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")
class MmapStore(Store, Streamable, Recordable):
"""Memory-mapped binary storage backend.
+32 -11
View File
@@ -4,27 +4,48 @@ Public API:
- ``attn_decode`` — single-query decode attention
- ``attn_prefill`` — multi-query prefill attention
- ``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):
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, True = keep)
scale: 0.0 = auto (1/sqrt(head_dim)); >0 = explicit
layout: "bhld" (default) or "blhd"
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
(blhd). Scale is always ``1/sqrt(head_dim)``.
Causal and mask can coexist — both are applied simultaneously.
Each wrapper dispatches to its compiled CUDA kernel (``astrai.extension.attn_*``)
when available, otherwise falls back to ``torch.nn.functional.scaled_dot_product_attention``.
Each wrapper calls its compiled CUDA kernel directly. Fallback to torch
SDPA is handled by the attention backend, not the wrapper functions.
"""
from astrai.extension.attention_backend import (
ATTN_BACKEND,
AttentionBackend,
CudaBackend,
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.ops import attention, attn_decode, attn_paged_decode, attn_prefill
from astrai.extension.rotary_backend import apply_rotary_emb
__all__ = [
"ATTN_BACKEND",
"AttentionBackend",
"CudaBackend",
"TorchNativeBackend",
"TensorLayout",
"attention",
"attn_backend",
"get_backend",
"attn_decode",
"attn_paged_decode",
"attn_prefill",
"attention",
"is_available",
"KERNEL_NAMES",
"apply_rotary_emb",
]
+401
View File
@@ -0,0 +1,401 @@
"""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
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.inference.core.cache import KVCache
_current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.ContextVar(
"attn_backend"
)
class ATTN_BACKEND(enum.Enum):
"""Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``."""
TORCH_NATIVE = "torch_native"
CUDA = "cuda"
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[ATTN_BACKEND, "AttentionBackend", type]):
"""Context manager to select an attention backend.
Mirrors ``torch.nn.attention.sdpa_kernel``. Accepts an
``ATTN_BACKEND`` enum value, a backend class, or a backend instance.
Examples::
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
...
with attn_backend(TorchNativeBackend):
...
with attn_backend(TorchNativeBackend()):
...
"""
if isinstance(backend, ATTN_BACKEND):
instance = _BACKEND_REGISTRY[backend]()
elif isinstance(backend, type) and issubclass(backend, AttentionBackend):
instance = backend()
elif isinstance(backend, AttentionBackend):
instance = backend
else:
raise TypeError(
f"expected ATTN_BACKEND, AttentionBackend type, 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 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)
q = q.permute(0, 2, 1, 3)
k = k.permute(0, 2, 1, 3)
v = v.permute(0, 2, 1, 3)
out = F.scaled_dot_product_attention(q, k, v, attn_mask, is_causal=is_causal)
out = out.permute(0, 2, 1, 3).contiguous().flatten(2)
return out
_default_backend = TorchNativeBackend()
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 = torch.arange(b + 1, dtype=torch.int32, device=q.device) * q_len
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)
_BACKEND_REGISTRY: dict[ATTN_BACKEND, type[AttentionBackend]] = {
ATTN_BACKEND.TORCH_NATIVE: TorchNativeBackend,
ATTN_BACKEND.CUDA: CudaBackend,
}
+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__)
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] = {}
_modules: dict[str, object] = {}
for _name in KERNEL_NAMES:
try:
_mod = importlib.import_module(f".{_name}", package=__package__)
_mod = importlib.import_module(f".lib.{_name}", package=__package__)
_available[_name] = True
_modules[_name] = _mod
except ImportError:
-298
View File
@@ -1,298 +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
)
def attention(
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:
"""Dispatch to decode or prefill attention based on the query length.
A query length of one is the decode case; longer queries use prefill.
The paged-cache decode path cannot be selected here because its page-table
arguments are not part of this interface.
"""
li = _parse_layout(layout)
if q.ndim not in (2, 3, 4) or k.ndim != q.ndim or v.ndim != q.ndim:
raise ValueError(
"q, k, and v must all have the same rank in {2, 3, 4}, "
f"got {q.ndim}D, {k.ndim}D, {v.ndim}D"
)
if k.shape != v.shape:
raise ValueError(
f"k and v must have the same shape, got {k.shape} and {v.shape}"
)
original_ndim = q.ndim
if original_ndim == 2:
# [L, D] -> [1, 1, L, D] or [1, L, 1, D]
q = q.unsqueeze(0).unsqueeze(1 if li == 0 else 2)
k = k.unsqueeze(0).unsqueeze(1 if li == 0 else 2)
v = v.unsqueeze(0).unsqueeze(1 if li == 0 else 2)
elif original_ndim == 3:
# [B, L, D] -> single-head 4D input.
q = q.unsqueeze(1 if li == 0 else 2)
k = k.unsqueeze(1 if li == 0 else 2)
v = v.unsqueeze(1 if li == 0 else 2)
q_len = q.size(2 if li == 0 else 1)
if q_len == 1:
out = attn_decode(q, k, v, mask, causal_offset, scale, layout)
else:
out = attn_prefill(q, k, v, mask, causal_offset, scale, layout)
if original_ndim == 2:
return out.squeeze(0).squeeze(0 if li == 0 else 1)
if original_ndim == 3:
return out.squeeze(1 if li == 0 else 2)
return out
+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 -35
View File
@@ -13,41 +13,63 @@ from typing import (
Type,
TypeVar,
Union,
get_args,
get_origin,
)
from typing import get_args as _get_args
from typing import get_origin as _get_origin
T = TypeVar("T")
def _resolve_type(
def _resolve_base_type(
arg: Union[Type, str, ForwardRef], factory_cls: type
) -> Optional[Type]:
"""Resolve a generic type-arg (str forward-ref, ForwardRef, or class)."""
if not isinstance(arg, (str, ForwardRef)):
"""Resolve the generic type-arg T to a concrete class.
- 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
name = arg if isinstance(arg, str) else arg.__forward_arg__
if name == factory_cls.__name__:
return factory_cls
if isinstance(arg, str):
name = arg
elif isinstance(arg, ForwardRef):
name = arg.__forward_arg__
else:
return None
mod = sys.modules.get(factory_cls.__module__)
if mod is 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]):
"""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]):
pass
Register components with the ``register`` decorator::
@MyFactory.register("custom")
class CustomComponent(MyBase):
...
@@ -64,13 +86,10 @@ class BaseFactory(ABC, Generic[T]):
def __init_subclass__(cls, **kwargs):
super().__init_subclass__(**kwargs)
for orig_base in getattr(cls, "__orig_bases__", ()):
if _get_origin(orig_base) is BaseFactory:
(arg,) = _get_args(orig_base)
if get_origin(orig_base) is BaseFactory:
(arg,) = get_args(orig_base)
cls._entries = {}
try:
cls._component_base = _resolve_type(arg, cls)
except Exception:
cls._component_base = None
cls._component_base = _resolve_base_type(arg, cls)
return
@classmethod
@@ -82,7 +101,7 @@ class BaseFactory(ABC, Generic[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:
raise ValueError(f"Component '{name}' is already registered")
cls._entries[name] = component_cls
@@ -95,12 +114,11 @@ class BaseFactory(ABC, Generic[T]):
"""Create a component instance by name, filtering kwargs to match
the component's ``__init__`` signature.
"""
entry = cls._entries.get(name)
if entry is None:
component_cls = cls._entries.get(name)
if component_cls is None:
raise ValueError(
f"Unknown component: '{name}'. Supported types: {sorted(cls._entries)}"
)
component_cls = entry
sig = inspect.signature(component_cls.__init__)
has_var_kwargs = any(
p.kind == inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values()
@@ -114,18 +132,6 @@ class BaseFactory(ABC, Generic[T]):
kwargs = {k: v for k, v in kwargs.items() if k in valid}
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
def get_component_class(cls, name: str) -> Type[T]:
"""Get the registered component class without instantiating it."""
+4 -14
View File
@@ -30,21 +30,16 @@ from astrai.inference.api.openai import OpenAIResponseBuilder
from astrai.inference.core import (
STOP,
Allocator,
CacheView,
ContiguousCache,
ContiguousCacheView,
Executor,
InferenceScheduler,
KVCache,
PageCache,
PageCacheView,
KVStorage,
PagePool,
PrefixCache,
Storage,
ReqToTokenPool,
Task,
TaskManager,
TaskStatus,
TaskTable,
page_hash,
)
from astrai.inference.engine import GenerationRequest, InferenceEngine
@@ -68,16 +63,11 @@ __all__ = [
"TaskManager",
"TaskStatus",
"Allocator",
"CacheView",
"KVCache",
"ContiguousCache",
"ContiguousCacheView",
"PageCache",
"PageCacheView",
"KVStorage",
"PagePool",
"PrefixCache",
"Storage",
"TaskTable",
"ReqToTokenPool",
"page_hash",
"sample",
"BaseSamplingStrategy",
+4
View File
@@ -110,6 +110,7 @@ def _create_engine(
device: str = "cuda",
dtype: torch.dtype = torch.bfloat16,
max_batch_size: int = 16,
max_seq_len: Optional[int] = None,
) -> InferenceEngine:
if not param_path.exists():
raise FileNotFoundError(f"Parameter directory not found: {param_path}")
@@ -123,6 +124,7 @@ def _create_engine(
model=model,
tokenizer=tokenizer,
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}")
return engine
@@ -186,6 +188,7 @@ def run_server(
device: str = "cuda",
dtype: torch.dtype = torch.bfloat16,
max_batch_size: int = 16,
max_seq_len: Optional[int] = None,
):
app = get_app()
app.state.server_config = {
@@ -193,6 +196,7 @@ def run_server(
"dtype": dtype,
"param_path": param_path,
"max_batch_size": max_batch_size,
"max_seq_len": max_seq_len,
}
uvicorn.run(
app,
+10 -15
View File
@@ -22,13 +22,10 @@ class BaseToolParser(ABC):
Maintains streaming state internally so that each call to :meth:`feed`
can diff against previously emitted content.
Parameters
----------
tools : list of dict, optional
Tool definitions from the request.
tool_choice : str
``"auto"`` / ``"required"`` / ``"none"`` or a named tool choice
dict.
Args:
tools (list of dict, optional): Tool definitions from the request.
tool_choice (str): ``"auto"`` / ``"required"`` / ``"none"`` or a named
tool choice dict.
"""
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.
Parameters
----------
body : str
The complete accumulated generated text so far.
current_token_ids : list of int, optional
All token IDs decoded into *body* (cumulative).
delta_token_ids : list of int, optional
Only the token IDs for this chunk.
Args:
body (str): The complete accumulated generated text so far.
current_token_ids (list of int, optional): All token IDs decoded
into *body* (cumulative).
delta_token_ids (list of int, optional): Only the token IDs for
this chunk.
"""
@abstractmethod
+4 -14
View File
@@ -2,16 +2,11 @@
from astrai.inference.core.cache import (
Allocator,
CacheView,
ContiguousCache,
ContiguousCacheView,
KVCache,
PageCache,
PageCacheView,
KVStorage,
PagePool,
PrefixCache,
Storage,
TaskTable,
ReqToTokenPool,
page_hash,
)
from astrai.inference.core.executor import Executor
@@ -20,16 +15,11 @@ from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
__all__ = [
"Allocator",
"CacheView",
"KVCache",
"ContiguousCache",
"ContiguousCacheView",
"PageCache",
"PageCacheView",
"KVStorage",
"PagePool",
"PrefixCache",
"Storage",
"TaskTable",
"ReqToTokenPool",
"page_hash",
"Executor",
"InferenceScheduler",
+322 -357
View File
@@ -1,7 +1,21 @@
"""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 PrefixCache (content addressing).
KVCache is a pure dataclass passed to the model for direct buffer access.
Two modes:
- contiguous (default): pre-allocated per-request blocks, no dynamic alloc
- paged: shared pool with on-demand allocation, prefix caching support
"""
import threading
from abc import ABC, abstractmethod
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
from torch import Tensor
@@ -108,418 +122,369 @@ class PrefixCache:
self._hash_to_page[h] = page_idx
class PagePool:
"""Orchestrates allocator (page management) and PrefixCache (content addressing)."""
class ReqToTokenPool:
"""Maps [req_idx, pos] -> physical token slot in KV storage.
def __init__(self, allocator: Allocator, prefix: PrefixCache):
self._alloc = allocator
self._prefix = prefix
self._alloc.on_evict = prefix.evict
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.
"""
@property
def allocator(self) -> Allocator:
return self._alloc
@property
def prefix(self) -> PrefixCache:
return self._prefix
def alloc(self) -> int:
return self._alloc.alloc()
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] = {}
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 set(self, task_id: str, page_table: List[int], cached: int):
def alloc(self, num_reqs: int) -> Optional[List[int]]:
with self._lock:
self._pages[task_id] = page_table
self._cached[task_id] = cached
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 get(self, task_id: str) -> List[int]:
def free(self, req_indices: List[int]):
with self._lock:
return self._pages.get(task_id, [])
self.free_slots.extend(req_indices)
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)
def write(self, indices, values):
self.req_to_token[indices] = values
class Storage:
"""KV-cache tensor storage with paged write/gather."""
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_pages: int,
page_size: int,
n_kv_heads: 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.size = size
self.k_buffer = torch.empty(
(n_layers, 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,
self.v_buffer = torch.empty(
(n_layers, 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 get_key_buffer(self, layer_id: int) -> Tensor:
return self.k_buffer[layer_id]
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
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
class CacheView(ABC):
"""Abstract view passed to attention layers for KV-cache I/O."""
@dataclass
class KVCache:
"""Pure data struct passed to model for KV cache I/O.
@abstractmethod
def write(self, layer_id: int, k: Tensor, v: Tensor): ...
The attention layer does raw buffer indexing — no methods, no abstraction.
@abstractmethod
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]: ...
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
class KVCache(ABC):
"""Abstract KV-cache facade for scheduler/executor."""
class PagePool:
"""Top-level KV cache manager.
@abstractmethod
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool: ...
Combines KVStorage + ReqToTokenPool + Allocator + PrefixCache.
@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."""
Args:
n_layers: Number of transformer layers.
n_kv_heads: Number of KV attention heads.
head_dim: Dimension per head.
max_batch_size: Maximum concurrent requests.
max_seq_len: Maximum sequence length per request.
device, dtype: Tensor device and dtype.
page_size: Page size for paged mode (1 = token-level).
n_tokens: Total token slots for paged mode. None = contiguous mode
(pre-allocates max_batch_size * max_seq_len).
"""
def __init__(
self,
n_layers: int,
n_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)
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
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
max_len = self._total_len
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_seq_len: int,
n_kv_heads: int,
head_dim: int,
device: torch.device,
dtype: torch.dtype,
page_size: int = 1,
n_tokens: Optional[int] = None,
):
self.page_size = page_size
self.max_batch_size = max_batch_size
self.max_seq_len = max_seq_len
self.k = torch.zeros(
n_layers,
max_batch_size,
max_seq_len,
n_kv_heads,
head_dim,
device=device,
dtype=dtype,
self.device = device
self.dtype = dtype
self.n_layers = n_layers
self.n_kv_heads = n_kv_heads
self.head_dim = head_dim
self.contiguous = n_tokens is None
if self.contiguous:
self.n_tokens = max_batch_size * max_seq_len
else:
self.n_tokens = n_tokens
self._storage = KVStorage(
self.n_tokens, n_layers, n_kv_heads, head_dim, device, dtype
)
self.v = torch.zeros(
n_layers,
max_batch_size,
max_seq_len,
n_kv_heads,
head_dim,
device=device,
dtype=dtype,
)
self._slot_len: Dict[int, int] = {}
self._task_slot: Dict[str, int] = {}
self._free_slots = list(range(max_batch_size))
self._device = device
self._req_pool = ReqToTokenPool(max_batch_size, max_seq_len, device)
if self.contiguous:
for i in range(max_batch_size):
self._req_pool.req_to_token[i] = torch.arange(
i * max_seq_len, (i + 1) * max_seq_len, device=device
)
self._alloc: Optional[Allocator] = None
self._prefix: Optional[PrefixCache] = None
else:
n_pages = self.n_tokens // page_size
self._alloc = Allocator(n_pages)
self._prefix = PrefixCache(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()
# ---- task lifecycle ----
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
slot = self._free_slots.pop(0)
self._task_slot[task_id] = slot
self._slot_len[slot] = 0
req_idx = req_slots[0]
self._task_req[task_id] = req_idx
if self.contiguous:
self._task_len[req_idx] = len(prompt_ids)
self._task_cached[task_id] = 0
return True
n_tokens_needed = len(prompt_ids)
cached = 0
if self._prefix is not None:
hits = self._prefix.lookup(prompt_ids)
cached = len(hits) * self.page_size
for p in hits:
self._alloc.inc_ref(p)
self._task_pages[task_id] = list(hits)
self._task_slots[task_id] = []
else:
self._task_pages[task_id] = []
self._task_slots[task_id] = []
remaining = n_tokens_needed - cached
if remaining > 0:
if self.page_size == 1:
slots = self._alloc_tokens(remaining)
if slots is None:
for p in self._task_pages[task_id]:
self._alloc.free(p)
self._req_pool.free([req_idx])
del self._task_req[task_id]
return False
self._task_slots[task_id] = slots
else:
n_new_pages = (remaining + self.page_size - 1) // self.page_size
new_pages = []
for _ in range(n_new_pages):
p = self._alloc.alloc()
if p < 0:
for hp in self._task_pages[task_id]:
self._alloc.free(hp)
for np_ in new_pages:
self._alloc.free(np_)
self._req_pool.free([req_idx])
del self._task_req[task_id]
return False
new_pages.append(p)
self._task_pages[task_id].extend(new_pages)
self._write_req_to_token(task_id, prompt_ids, cached)
self._task_len[req_idx] = len(prompt_ids)
self._task_cached[task_id] = cached
return True
def task_free(self, task_id: str):
slot = self._task_slot.pop(task_id, None)
if slot is not None:
self._slot_len.pop(slot, None)
self._free_slots.append(slot)
req_idx = self._task_req.pop(task_id, None)
if req_idx is None:
return
self._task_len.pop(req_idx, None)
self._task_cached.pop(task_id, None)
if not self.contiguous:
if self._prefix is not None:
for p in self._task_pages.get(task_id, []):
keep = self._prefix.has_page(p)
self._alloc.free(p, keep_cached=keep)
if not keep:
self._prefix.evict(p)
else:
for p in self._task_pages.get(task_id, []):
self._alloc.free(p)
self._task_pages.pop(task_id, None)
self._task_slots.pop(task_id, None)
self._req_pool.free([req_idx])
def task_extend(self, task_id: str, pos: int) -> bool:
return pos < self.max_seq_len
req_idx = self._task_req.get(task_id)
if req_idx is None:
return False
if self.contiguous:
return pos < self.max_seq_len
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]
else:
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
self._task_len[req_idx] = pos + 1
return True
def task_cached(self, task_id: str) -> int:
slot = self._task_slot.get(task_id)
if slot is None:
return 0
return self._slot_len.get(slot, 0)
return self._task_cached.get(task_id, 0)
def task_record_hashes(
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
):
if self._prefix is None or self.contiguous:
return
pages = self._task_pages.get(task_id, [])
full_pages = len(prompt_ids) // self.page_size
for i in range(start_logical_page, min(full_pages, len(pages))):
self._prefix.record(pages[i], prompt_ids, i)
# ---- bind for forward ----
def bind_tasks(
self,
task_ids: List[str],
total_len: int,
seq_lens: List[int],
device: torch.device,
write_positions: Optional[Tensor] = None,
) -> ContiguousCacheView:
slots = [self._task_slot[tid] for tid in task_ids]
batch_indices = torch.tensor(slots, dtype=torch.long, device=device)
for slot in slots:
if total_len > self._slot_len.get(slot, 0):
self._slot_len[slot] = total_len
return ContiguousCacheView(
self, batch_indices, total_len, write_positions=write_positions
start_pos: Optional[int] = None,
) -> KVCache:
req_indices = [self._task_req[tid] for tid in task_ids]
req_pool_indices = torch.tensor(req_indices, dtype=torch.long, device=device)
seq_lens_t = torch.tensor(seq_lens, dtype=torch.long, device=device)
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
]
else:
write_pos = seq_lens_t - 1
out_cache_loc = self._req_pool.req_to_token[
req_pool_indices, write_pos
].unsqueeze(-1)
kv_indptr = torch.zeros(len(seq_lens) + 1, dtype=torch.int32, device=device)
kv_indptr[1:] = seq_lens_t.cumsum(0).to(torch.int32)
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,
)
# ---- 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
+33 -25
View File
@@ -3,7 +3,7 @@ from typing import List, Optional
import torch
from astrai.inference.core.cache import KVCache
from astrai.inference.core.cache import PagePool
from astrai.inference.core.task import Task
from astrai.inference.sample import sample
from astrai.model.automodel import AutoModel
@@ -19,7 +19,7 @@ class Executor:
self,
model: AutoModel,
tokenizer: AutoTokenizer,
kv_cache: KVCache,
kv_cache: PagePool,
device: Optional[str] = None,
dtype: Optional[torch.dtype] = None,
):
@@ -57,7 +57,9 @@ class Executor:
input_ids,
input_mask=input_mask,
position_ids=position_ids,
paged_cache=self.kv_cache.bind_tasks(task_ids, prompt_len, self.device),
kv_cache=self.kv_cache.bind_tasks(
task_ids, [prompt_len] * batch_sz, self.device, start_pos=start_pos
),
)
def execute_decode(
@@ -103,36 +105,42 @@ class Executor:
[t.frequency_penalty for t in tasks], device=self.device
)
history_lists = []
history_lens = []
for t in tasks:
window = t.rep_window
prompt_part = t.prompt_ids[-window:]
ids = prompt_part + t.output_ids
history_lists.append(ids)
history_lens.append(len(ids))
has_freq = bool((freq_penalties != 0).any())
if has_freq:
history_lists = []
history_lens = []
for t in tasks:
window = t.rep_window
prompt_part = t.prompt_ids[-window:]
ids = prompt_part + t.output_ids
history_lists.append(ids)
history_lens.append(len(ids))
max_len = max(history_lens) if history_lens else 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, h in enumerate(history_lists):
L = history_lens[i]
padded_ids[i, :L] = torch.as_tensor(h, dtype=torch.long, device=self.device)
padded_mask[i, :L] = True
max_len = max(history_lens) if history_lens else 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, h in enumerate(history_lists):
L = history_lens[i]
padded_ids[i, :L] = torch.as_tensor(
h, dtype=torch.long, device=self.device
)
padded_mask[i, :L] = True
else:
padded_ids = None
padded_mask = None
with torch.inference_mode():
outputs = self.model(
input_ids.unsqueeze(1),
input_mask=input_mask,
paged_cache=self.kv_cache.bind_tasks(
kv_cache=self.kv_cache.bind_tasks(
task_ids,
total_len,
[t.next_pos + 1 for t in tasks],
self.device,
write_positions=position_ids,
),
position_ids=position_ids.unsqueeze(1),
)
+15 -15
View File
@@ -5,7 +5,7 @@ from typing import Any, Dict, List, Optional, Tuple
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.task import STOP, Task, TaskManager, TaskStatus
from astrai.model.automodel import AutoModel
@@ -23,10 +23,9 @@ class InferenceScheduler:
tokenizer: AutoTokenizer,
max_batch_size: int = 16,
max_seq_len: Optional[int] = None,
max_prompt_len: int = 2048,
device: Optional[str] = None,
dtype: Optional[torch.dtype] = None,
cache: Optional[KVCache] = None,
cache: Optional[PagePool] = None,
):
config = model.config
@@ -47,21 +46,20 @@ class InferenceScheduler:
if cache is not None:
self._cache = cache
else:
self._cache = ContiguousCache(
config.num_hidden_layers,
max_batch_size,
self.max_seq_len,
config.num_key_value_heads,
head_dim,
self.device,
self.dtype,
self._cache = PagePool(
n_layers=config.num_hidden_layers,
n_kv_heads=config.num_key_value_heads,
head_dim=head_dim,
max_batch_size=max_batch_size,
max_seq_len=self.max_seq_len,
device=self.device,
dtype=self.dtype,
)
self._task_mgr = TaskManager(
tokenizer=tokenizer,
max_batch_size=max_batch_size,
max_seq_len=self.max_seq_len,
max_prompt_len=max_prompt_len,
)
self._executor = Executor(
@@ -111,9 +109,11 @@ class InferenceScheduler:
self._task_mgr.wait_for_tasks(timeout=1.0)
continue
active = self._task_mgr.get_active_tasks()
to_prefill = [
t
for t in self._task_mgr.get_active_tasks()
for t in active
if t.output_tokens == 0
and cache.task_cached(t.task_id) < len(t.prompt_ids)
]
@@ -139,10 +139,10 @@ class InferenceScheduler:
t.task_id, t.prompt_ids, start_logical_page
)
decode_tasks = self._task_mgr.get_active_tasks()
decode_tasks = active
valid: List[Task] = []
for t in sorted(decode_tasks, key=lambda t: t.task_id):
for t in decode_tasks:
if cache.task_extend(t.task_id, t.next_pos):
valid.append(t)
else:
+22 -36
View File
@@ -6,6 +6,8 @@ from collections import deque
from enum import Enum
from typing import Any, Callable, Deque, Dict, List, Optional
from tokenizers.decoders import DecodeStream
from astrai.tokenize.tokenizer import AutoTokenizer
logger = logging.getLogger(__name__)
@@ -14,37 +16,30 @@ STOP = object()
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,
smart quotes) across multiple tokens. Decoding such a token in
isolation produces U+FFFD (replacement char). This decoder
accumulates token IDs and only emits text once the trailing
characters are complete, buffering incomplete multi-byte sequences
until the next token arrives.
Delegates to the Rust-native streaming decoder which maintains an
O(1) bounded token buffer internally (via prefix drain), avoiding
the O(n²) cost of re-decoding the full history on each step.
Multi-byte UTF-8 sequences split across token boundaries are
buffered until complete; ``push`` returns "" while the trailing
sequence is still incomplete.
"""
__slots__ = ("_tokenizer", "_ids", "_emitted")
__slots__ = ("_stream", "_tok")
def __init__(self, tokenizer: AutoTokenizer):
self._tokenizer = tokenizer
self._ids: List[int] = []
self._emitted: str = ""
self._tok = tokenizer._tokenizer
self._stream = DecodeStream(skip_special_tokens=True)
def push(self, token_id: int) -> str:
"""Append a token ID and return newly completed text.
Returns "" while a multi-byte character is still incomplete.
"""
self._ids.append(token_id)
full = self._tokenizer.decode(self._ids, skip_special_tokens=True)
if full.endswith("\ufffd"):
return ""
if len(full) > len(self._emitted):
diff = full[len(self._emitted) :]
self._emitted = full
return diff
return ""
chunk = self._stream.step(self._tok, token_id)
return chunk or ""
class TaskStatus(Enum):
@@ -101,18 +96,11 @@ class Task:
def flush_remaining(self, tokenizer: AutoTokenizer) -> str:
"""Emit any text still buffered in the decoder.
Called when generation terminates (max_tokens reached, stop
sequence, or external removal) to avoid dropping a final
incomplete-looking fragment that is actually complete when
adjacent to the stop token.
With the Rust-native DecodeStream, the stream is always in a
correct state — any completed text was already emitted by the
last ``push``. A trailing incomplete multi-byte sequence has no
valid text to emit, so this is a no-op.
"""
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 ""
@property
@@ -135,12 +123,10 @@ class TaskManager:
tokenizer: AutoTokenizer,
max_batch_size: int = 16,
max_seq_len: int = 8192,
max_prompt_len: int = 512,
):
self.tokenizer = tokenizer
self.max_batch_size = max_batch_size
self.max_seq_len = max_seq_len
self.max_prompt_len = max_prompt_len
self.waiting_queue: Deque[Task] = deque()
self.active_tasks: List[Task] = []
@@ -165,10 +151,10 @@ class TaskManager:
) -> str:
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
prompt_ids = self.tokenizer.encode(prompt)
if len(prompt_ids) > self.max_prompt_len:
prompt_ids = prompt_ids[-self.max_prompt_len :]
if len(prompt_ids) > self.max_seq_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:
stream_callback(STOP)
return task_id
+2 -5
View File
@@ -8,7 +8,7 @@ from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Tuple,
import torch
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.task import STOP
from astrai.tokenize import AutoTokenizer
@@ -111,9 +111,7 @@ class InferenceEngine:
tokenizer: AutoTokenizer,
max_batch_size: int = 1,
max_seq_len: Optional[int] = None,
max_prompt_len: int = 2048,
page_size: int = 128,
cache: Optional[KVCache] = None,
cache: Optional[PagePool] = None,
):
self.model = model
self.tokenizer = tokenizer
@@ -122,7 +120,6 @@ class InferenceEngine:
tokenizer=self.tokenizer,
max_batch_size=max_batch_size,
max_seq_len=max_seq_len,
max_prompt_len=max_prompt_len,
cache=cache,
)
+37 -8
View File
@@ -343,6 +343,10 @@ def sample(
When **temperature** is exactly 0 (scalar or single-element tensor)
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:
logits: Raw logits ``[batch, vocab_size]``.
frequency_penalty: Penalty per occurrence for repeated tokens
@@ -359,14 +363,39 @@ def sample(
``True`` — a ``(token_ids, chosen_logprobs)`` tuple where
``chosen_logprobs`` has shape ``[batch]``.
"""
return SamplingPipeline(
[
TemperatureStrategy(temperature),
TopKStrategy(top_k),
TopPStrategy(top_p),
FrequencyPenaltyStrategy(frequency_penalty),
]
).sample(
greedy = (
(
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),
TopKStrategy(top_k),
TopPStrategy(top_p),
]
if has_freq:
strategies.append(FrequencyPenaltyStrategy(frequency_penalty))
return SamplingPipeline(strategies).sample(
logits,
filter_value=filter_value,
input_ids=input_ids,
+2 -1
View File
@@ -9,7 +9,7 @@ from astrai.model.components.lora import (
merge_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.encoder import EmbeddingEncoder
from astrai.model.transformer import AutoRegressiveLM
@@ -19,6 +19,7 @@ __all__ = [
"Linear",
"RMSNorm",
"MLP",
"DeepSeekMoE",
"GQA",
"DecoderBlock",
# Models
+7 -6
View File
@@ -40,11 +40,12 @@ def _disable_random_init(enable: bool = True):
setattr(nn.init, n, fn)
class AutoModel(BaseFactory["AutoModel"], nn.Module):
"""
Autoregressive language model base class.
Provides model loading/saving, registration, and generation.
"""
class ModelFactory(BaseFactory[nn.Module]):
"""Pure factory for model dispatch, separated from nn.Module state."""
class AutoModel(nn.Module):
"""Model base class with loading/saving and generation."""
def __init__(self, config: BaseModelConfig):
super().__init__()
@@ -68,7 +69,7 @@ class AutoModel(BaseFactory["AutoModel"], nn.Module):
config = ConfigFactory.load(raw)
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):
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.embedding import Embedding
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.rope import (
RotaryEmbedding,
apply_rotary_emb,
get_rotary_emb,
)
@@ -14,6 +14,7 @@ __all__ = [
"Linear",
"RMSNorm",
"MLP",
"DeepSeekMoE",
"Embedding",
"GQA",
"MLA",
@@ -21,5 +22,4 @@ __all__ = [
"RotaryEmbedding",
"apply_rotary_emb",
"get_rotary_emb",
"repeat_kv",
]
+7 -40
View File
@@ -5,22 +5,12 @@ import torch.nn as nn
import torch.nn.functional as F
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.inference.core.cache import CacheView
from astrai.inference.core.cache import KVCache
from astrai.model.components.linear import Linear
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]):
@@ -75,7 +65,7 @@ class GQA(nn.Module):
x: Tensor,
rotary_emb: Tensor,
attn_mask: Tensor = None,
paged_cache: Optional[CacheView] = None,
kv_cache: Optional[KVCache] = None,
is_causal: bool = False,
) -> Tensor:
q = self._split_heads(self.q_proj(x), self.n_heads)
@@ -86,19 +76,7 @@ class GQA(nn.Module):
if self.use_qk_norm:
q, k = self.q_norm(q), self.k_norm(k)
if paged_cache is not None:
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)
)
sdqa_out = attention(q, k, v, kv_cache, self.layer_id, attn_mask, is_causal)
if self.use_gated_attention:
sdqa_out = sdqa_out * F.sigmoid(self.gate(x))
@@ -161,7 +139,7 @@ class MLA(nn.Module):
x: Tensor,
rotary_emb: Tensor,
attn_mask: Tensor = None,
paged_cache: Optional[CacheView] = None,
kv_cache: Optional[KVCache] = None,
is_causal: bool = False,
) -> Tensor:
bsz, seq_len, _ = x.size()
@@ -193,18 +171,7 @@ class MLA(nn.Module):
q = self.q_norm(q)
k = self.k_norm(k)
if paged_cache is not None:
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)
attn_out = attention(q, k, v, kv_cache, self.layer_id, attn_mask, is_causal)
if self.use_gated_attention:
attn_out = attn_out * F.sigmoid(self.gate(x))
+28 -8
View File
@@ -1,15 +1,20 @@
from dataclasses import asdict
from typing import Optional
from typing import Optional, TypedDict
import torch.nn as nn
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.mlp import FFNFactory
from astrai.model.components.norm import RMSNorm
class DecoderOutput(TypedDict):
hidden_states: Tensor
aux_loss: Optional[Tensor]
class DecoderBlock(nn.Module):
def __init__(self, config, layer_id: int):
super().__init__()
@@ -26,24 +31,39 @@ class DecoderBlock(nn.Module):
self.attention = AttnFactory.create(config.attn_type, **cfg, layer_id=layer_id)
self.input_norm = RMSNorm(config.hidden_size, config.rms_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(
self,
x: Tensor,
rotary_emb: Tensor,
attention_mask: Optional[Tensor] = None,
paged_cache: Optional[CacheView] = None,
kv_cache: Optional[KVCache] = None,
is_causal: bool = False,
) -> Tensor:
) -> DecoderOutput:
attn_output = self.attention(
self.input_norm(x),
rotary_emb,
attention_mask,
paged_cache,
kv_cache,
is_causal,
)
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"]}
+2 -1
View File
@@ -1,11 +1,12 @@
import logging
from dataclasses import asdict, dataclass
from dataclasses import asdict
from pathlib import Path
from typing import Optional, Set
import torch
import torch.nn as nn
import torch.nn.functional as F
from pydantic.dataclasses import dataclass
from astrai.model.components.linear import Linear
from astrai.serialization import (
+54 -13
View File
@@ -1,3 +1,5 @@
from typing import Optional, TypedDict
import torch
import torch.nn as nn
import torch.nn.functional as F
@@ -11,6 +13,16 @@ class FFNFactory(BaseFactory[nn.Module]):
pass
class FFNOutput(TypedDict):
hidden_states: Tensor
aux_loss: Optional[Tensor]
class RoutedOutput(TypedDict):
hidden_states: Tensor
aux_loss: Optional[Tensor]
@FFNFactory.register("mlp")
class MLP(nn.Module):
def __init__(self, dim: int, dim_ffn: int, down_init_std: float = 0.02):
@@ -19,10 +31,10 @@ class MLP(nn.Module):
self.gate = Linear(dim, dim_ffn)
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))
out = self.down(gated)
return out
return {"hidden_states": out, "aux_loss": None}
@FFNFactory.register("moe")
@@ -36,6 +48,9 @@ class DeepSeekMoE(nn.Module):
n_activated_experts: int = 2,
topk_method: str = "greedy",
n_layers: int = 1,
moe_intermediate_size: Optional[int] = None,
shared_expert_intermediate_size: Optional[int] = None,
norm_topk_prob: bool = True,
):
super().__init__()
self.dim = dim
@@ -43,6 +58,16 @@ class DeepSeekMoE(nn.Module):
self.n_shared_experts = n_shared_experts
self.n_activated_experts = n_activated_experts
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)
moe_scale = 1 / max(n_shared_experts, 1) + 1 / n_activated_experts
@@ -50,33 +75,37 @@ class DeepSeekMoE(nn.Module):
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)
]
)
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)
]
)
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
x_flat = x.view(-1, dim)
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)
return out
out = (shared_out + routed_output["hidden_states"]).view(bsz, seq_len, dim)
return {"hidden_states": out, "aux_loss": routed_output["aux_loss"]}
def _shared_forward(self, x: Tensor) -> Tensor:
if self.n_shared_experts == 0:
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
K = self.n_activated_experts
@@ -84,7 +113,17 @@ class DeepSeekMoE(nn.Module):
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_weights / topk_weights.sum(dim=-1, keepdim=True)
if self.norm_topk_prob:
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
aux_loss = None
if include_aux_loss:
expert_load = F.one_hot(
topk_indices, num_classes=self.n_routed_experts
).float()
expert_load = expert_load.mean(dim=(0, 1))
router_prob = router_probs.float().mean(dim=0)
aux_loss = self.n_routed_experts * (expert_load * router_prob).sum()
output = torch.zeros(N, D, device=x.device, dtype=x.dtype)
for expert_idx in range(self.n_routed_experts):
@@ -92,9 +131,11 @@ class DeepSeekMoE(nn.Module):
token_idx, k_idx = expert_mask.nonzero(as_tuple=True)
if token_idx.numel() == 0:
continue
expert = self.routed_experts[expert_idx]
expert_input = x[token_idx]
expert_output = self.routed_experts[expert_idx](expert_input)
expert_output = expert(expert_input)["hidden_states"]
weights = topk_weights[token_idx, k_idx].unsqueeze(-1)
output.index_add_(0, token_idx, expert_output * weights)
return output
return {"hidden_states": output, "aux_loss": aux_loss}
+17 -15
View File
@@ -11,28 +11,23 @@ def get_rotary_emb(
base: float = 10000,
device: Optional[torch.device] = None,
) -> 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)
t = torch.arange(0, max_len, dtype=torch.float64, device=device)
freqs = torch.outer(t, theta).float()
cos = torch.cos(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:
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):
def __init__(
self,
@@ -56,16 +51,23 @@ class RotaryEmbedding(nn.Module):
self._set_rotary_buffer(self.max_len)
def _set_rotary_buffer(self, max_len: int):
rotary_emb = get_rotary_emb(self.dim, max_len, self.base)
freqs_cis = torch.view_as_real(rotary_emb)
freqs_cis = get_rotary_emb(self.dim, max_len, self.base)
self.register_buffer("freqs_cis", freqs_cis, persistent=False)
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:
position_ids = (
torch.arange(x.size(1), device=x.device)
.unsqueeze(0)
.expand(x.size(0), -1)
)
position_freq_cis = self.freqs_cis[position_ids].float()
return torch.view_as_complex(position_freq_cis)
return self.freqs_cis[position_ids].float()
+3 -3
View File
@@ -5,7 +5,7 @@ import torch.nn as nn
from torch import Tensor
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.embedding import Embedding
from astrai.model.components.norm import RMSNorm
@@ -13,7 +13,7 @@ from astrai.model.components.rope import RotaryEmbedding
from astrai.model.transformer import process_attention_mask
@AutoModel.register("embedding")
@ModelFactory.register("embedding")
class EmbeddingEncoder(AutoModel):
def __init__(self, config: EncoderConfig):
super().__init__(config)
@@ -70,7 +70,7 @@ class EmbeddingEncoder(AutoModel):
attn_mask = process_attention_mask(input_mask)
for layer in self.layers:
x = layer(x, rotary_emb, attn_mask)
x = layer(x, rotary_emb, attn_mask)["hidden_states"]
hidden_states = self.norm(x)
+19 -6
View File
@@ -5,8 +5,8 @@ import torch.nn as nn
from torch import Tensor
from astrai.config.model_config import AutoRegressiveLMConfig
from astrai.inference.core.cache import CacheView
from astrai.model.automodel import AutoModel
from astrai.inference.core.cache import KVCache
from astrai.model.automodel import AutoModel, ModelFactory
from astrai.model.components.decoder_block import DecoderBlock
from astrai.model.components.embedding import Embedding
from astrai.model.components.linear import Linear
@@ -26,7 +26,7 @@ def process_attention_mask(
return input_mask
@AutoModel.register("autoregressive_lm")
@ModelFactory.register("autoregressive_lm")
class AutoRegressiveLM(AutoModel):
"""Autoregressive language model with paged KV cache."""
@@ -103,7 +103,7 @@ class AutoRegressiveLM(AutoModel):
self,
input_ids: Tensor,
input_mask: Optional[Tensor] = None,
paged_cache: Optional[CacheView] = None,
kv_cache: Optional[KVCache] = None,
position_ids: Optional[Tensor] = None,
) -> Dict[str, Tensor]:
assert input_ids.ndim == 2
@@ -113,10 +113,23 @@ class AutoRegressiveLM(AutoModel):
attn_mask = process_attention_mask(input_mask)
use_sdpa_causal_mask = attn_mask is None
aux_losses = []
for layer in self.layers:
x = layer(x, rotary_emb, attn_mask, paged_cache, use_sdpa_causal_mask)
layer_output = layer(
x,
rotary_emb,
attn_mask,
kv_cache,
use_sdpa_causal_mask,
)
x = layer_output["hidden_states"]
if layer_output["aux_loss"] is not None:
aux_losses.append(layer_output["aux_loss"])
hidden_states = self.norm(x)
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()
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 -5
View File
@@ -4,12 +4,12 @@ from astrai.parallel.executor import (
BaseExecutor,
DDPExecutor,
ExecutorFactory,
FSDP2Executor,
FSDPExecutor,
GradientState,
NoneExecutor,
broadcast_state_dict,
create_ref_model,
)
from astrai.parallel.module import ColumnParallelLinear, RowParallelLinear
from astrai.parallel.setup import (
get_current_device,
get_rank,
@@ -26,8 +26,6 @@ __all__ = [
"only_on_rank",
"setup_parallel",
"spawn_parallel_fn",
"RowParallelLinear",
"ColumnParallelLinear",
"ExecutorFactory",
"BaseExecutor",
"GradientState",
@@ -36,5 +34,6 @@ __all__ = [
"NoneExecutor",
"DDPExecutor",
"FSDPExecutor",
"FSDP2Executor",
"create_ref_model",
"broadcast_state_dict",
]
+120 -99
View File
@@ -4,18 +4,15 @@ import contextlib
import logging
import os
from contextlib import contextmanager
from typing import Any, Callable, Optional, Tuple
from typing import Any, Callable, Dict, Optional, Tuple
import torch
import torch.distributed as dist
import torch.nn as nn
from torch.distributed.fsdp import (
FSDPModule,
FullStateDictConfig,
StateDictType,
fully_shard,
)
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.tensor import DTensor
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.optim import Optimizer
@@ -27,6 +24,82 @@ from astrai.parallel.setup import get_rank, get_world_size
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:
def __init__(self, grad_accum_steps: int = 1):
self.num_steps = max(grad_accum_steps, 1)
@@ -95,11 +168,14 @@ class BaseExecutor:
optimizer_fn: Optional[Callable[[nn.Module], Optimizer]] = None,
scheduler_fn: Optional[Callable[[Optimizer], LRScheduler]] = None,
before_wrap: Optional[Callable[[nn.Module], nn.Module]] = None,
after_wrap: Optional[Callable[[nn.Module], nn.Module]] = None,
) -> 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)
if after_wrap is not None:
model = after_wrap(model)
optimizer = None
scheduler = None
if optimizer_fn is not None:
@@ -238,88 +314,11 @@ class DDPExecutor(BaseExecutor):
@ExecutorFactory.register("fsdp")
class FSDPExecutor(BaseExecutor):
def __init__(
self,
grad_accum_steps: int = 1,
process_group=None,
sharding_strategy=None,
cpu_offload=None,
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)
self._fsdp_kwargs = {
k: v
for k, v in dict(
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:
if not self.use_distributed:
logger.warning("FSDP backend selected but world_size=1, model not wrapped")
return model
self._original_model = model
device_id = torch.device("cuda", get_rank())
model = FSDP(model, device_id=device_id, **self._fsdp_kwargs)
logger.info("Model wrapped with FSDP (world_size=%d)", get_world_size())
return model
def _no_sync(self, model: nn.Module):
if isinstance(model, FSDP):
return model.no_sync()
return contextlib.nullcontext()
def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
if isinstance(model, FSDP) and self.use_distributed:
total_norm = model.clip_grad_norm_(max_norm)
if isinstance(total_norm, torch.Tensor):
return total_norm.item()
return total_norm
return super().clip_grad_norm(model, max_norm)
def unwrap_model(self, model: nn.Module):
if isinstance(model, FSDP) and 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()
@ExecutorFactory.register("fsdp2")
class FSDP2Executor(BaseExecutor):
"""FSDP2 executor using `torch.distributed.fsdp.fully_shard` (per-module API).
"""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
FSDP2's dynamic ``__class__`` assignment fail at the CPython level.
``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.
"""
@@ -329,7 +328,7 @@ class FSDP2Executor(BaseExecutor):
grad_accum_steps: int = 1,
mesh: Optional[Any] = None,
mp_policy: Optional[Any] = None,
reshard_after_forward: bool = True,
reshard_after_forward: bool = False,
):
super().__init__(grad_accum_steps=grad_accum_steps)
self._mesh = mesh
@@ -338,7 +337,7 @@ class FSDP2Executor(BaseExecutor):
def _prepare_model(self, model: nn.Module) -> nn.Module:
if not self.use_distributed:
logger.warning("FSDP2 backend selected but world_size=1, model not wrapped")
logger.warning("FSDP backend selected but world_size=1, model not wrapped")
return model
kwargs = dict(
@@ -356,7 +355,7 @@ class FSDP2Executor(BaseExecutor):
fully_shard(child, **kwargs)
logger.info(
"FSDP2 wrapping applied to %d direct children (root skipped for ABC compat)",
"FSDP wrapping applied to %d direct children (root skipped for ABC compat)",
len(list(model.children())),
)
return model
@@ -376,32 +375,54 @@ class FSDP2Executor(BaseExecutor):
yield
def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
if self.use_distributed:
total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)
if isinstance(total_norm, torch.Tensor):
return total_norm.item()
return total_norm
return super().clip_grad_norm(model, max_norm)
if not self.use_distributed:
return super().clip_grad_norm(model, max_norm)
# FSDP params are DTensors (sharded across ranks).
# 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()
def unwrap_model(self, model: nn.Module):
if not self.use_distributed:
return model.state_dict()
if get_rank() != 0:
return None
# 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 = {
k: (v.full_tensor() if isinstance(v, DTensor) else v)
for k, v in state_dict.items()
}
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)
+1 -1
View File
@@ -12,7 +12,7 @@ import torch
import torch.distributed as dist
import torch.multiprocessing as mp
from astrai.parallel.signal_handler import install_early_signal_handlers
from astrai.signal_handler import install_early_signal_handlers
logger = logging.getLogger(__name__)
+1 -1
View File
@@ -1,7 +1,7 @@
"""Config-driven JSONL preprocessing pipeline.
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,
dispatched by configuration keys.
+2 -21
View File
@@ -1,7 +1,7 @@
"""Storage writer strategies for pipeline output.
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
by ``output.storage_format``.
"""
@@ -15,7 +15,7 @@ from typing import Dict, List
import torch
from astrai.factory import BaseFactory
from astrai.serialization import save_bin, save_h5
from astrai.serialization import save_bin
logger = logging.getLogger(__name__)
@@ -54,22 +54,3 @@ class BinWriter(StoreWriter):
exc_info=True,
)
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 (
load_bin,
load_bin_offsets,
load_h5,
save_bin,
save_h5,
)
__all__ = [
@@ -39,7 +37,5 @@ __all__ = [
"save_torch",
"load_bin",
"load_bin_offsets",
"load_h5",
"save_bin",
"save_h5",
]
+4 -45
View File
@@ -1,55 +1,14 @@
"""Dataset storage serialization helpers (HDF5 / memory-mapped binary)."""
"""Dataset storage serialization helpers (memory-mapped binary)."""
import json
import os
from pathlib import Path
from typing import Any, Dict, List, Optional
import h5py
import numpy as np
import torch
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(
file_path: str,
tensor_group: Dict[str, List[Tensor]],
@@ -65,7 +24,7 @@ def save_bin(
offsets, preserving backward compatibility.
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)
record_keys = set(record_keys or [])
@@ -74,7 +33,7 @@ def save_bin(
if tensors and isinstance(tensors[0], list):
raise ValueError(
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)
entry: Dict[str, Any] = {
@@ -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),
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:
meta = json.load(f)
+18 -3
View File
@@ -38,12 +38,27 @@ class ChatTemplate:
The compiled :class:`~jinja2.Template` holds a dynamically-generated
``root`` render function whose ``__module__`` is ``None``; under
``pickle`` it falls back to ``__main__`` and breaks ``spawn``-based
multiprocessing. By deferring compilation to first access, the
default pickle protocol serialises only ``template_str``; each
worker rebuilds the cache on first render.
multiprocessing. :meth:`__getstate__` drops the cached template so
that pickle serialises only ``template_str``; each worker rebuilds
the cache on first render.
"""
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
def from_string(
cls,
-13
View File
@@ -20,8 +20,6 @@ Messages = List[Message]
class AutoTokenizer:
"""Base tokenizer class with automatic loading support"""
TOKENIZER_CLASSES = {} # Registry for auto-loading
def __init__(
self,
path: Optional[Union[str, Path]] = None,
@@ -108,17 +106,6 @@ class AutoTokenizer:
with open(save_path / "tokenizer_config.json", "w", encoding="utf-8") as f:
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(
self,
tokens: Union[str, List[str]],
+52
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()
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):
return ctx.loss
@@ -36,3 +81,10 @@ def ctx_get_val_loss(ctx):
def ctx_get_grad_norm(ctx):
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
+101 -40
View File
@@ -1,7 +1,7 @@
"""Training strategy implementations with factory pattern."""
from abc import ABC, abstractmethod
from typing import Callable, Dict, Union
from typing import Callable, Dict, Optional, TypedDict, Union
import torch
import torch.nn as nn
@@ -9,18 +9,18 @@ import torch.nn.functional as F
from torch import Tensor
from astrai.factory import BaseFactory
from astrai.parallel.executor import broadcast_state_dict
from astrai.trainer.rollout import RolloutResult
def create_ref_model(
model_fn: Callable[[], nn.Module], state_dict: Dict[str, Tensor]
) -> nn.Module:
"""Create a frozen reference model from model_fn + full state dict."""
ref_model = model_fn()
ref_model.load_state_dict(state_dict)
ref_model.requires_grad_(False)
ref_model.eval()
return ref_model
class LossOutput(TypedDict):
loss: Tensor
metrics: Dict[str, float]
class LogprobsOutput(TypedDict):
logprobs: Tensor
aux_loss: Optional[Tensor]
def move_to_device(batch: Dict[str, Tensor], device: str) -> Dict[str, Tensor]:
@@ -34,7 +34,7 @@ def get_logprobs(
attn_mask: Tensor,
loss_mask: Tensor,
reduction: str,
) -> Tensor:
) -> LogprobsOutput:
"""Compute token-wise log probabilities from model outputs.
Args:
@@ -56,10 +56,11 @@ def get_logprobs(
shifted_input_ids = input_ids[:, 1:]
shifted_loss_mask = loss_mask[:, 1:]
logits = model(
outputs = model(
input_ids[:, :-1],
attn_mask[:, :, :-1, :-1] if attn_mask.dim() == 4 else attn_mask[:, :-1],
)["logits"]
)
logits = outputs["logits"]
log_probs = torch.log_softmax(logits.float(), dim=-1)
token_logprobs = torch.gather(
@@ -67,13 +68,14 @@ def get_logprobs(
).squeeze(-1)
if reduction == "mean":
return (token_logprobs * shifted_loss_mask).sum(dim=-1) / shifted_loss_mask.sum(
logprobs = (token_logprobs * shifted_loss_mask).sum(
dim=-1
).clamp(min=1.0)
) / shifted_loss_mask.sum(dim=-1).clamp(min=1.0)
elif reduction == "sum":
return (token_logprobs * shifted_loss_mask).sum(dim=-1)
logprobs = (token_logprobs * shifted_loss_mask).sum(dim=-1)
else:
return token_logprobs * shifted_loss_mask
logprobs = token_logprobs * shifted_loss_mask
return {"logprobs": logprobs, "aux_loss": outputs.get("aux_loss")}
def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
@@ -112,6 +114,7 @@ class BaseStrategy(ABC):
self.model = model
self.device = device
self.executor = kwargs.pop("executor", None)
self.moe_aux_loss_coef = kwargs.pop("moe_aux_loss_coef", 0.01)
self.extra_kwargs = kwargs
self._rollout_runner = None
@@ -127,6 +130,33 @@ class BaseStrategy(ABC):
"""
raise NotImplementedError
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
return self._normalize_output(self.compute_loss(batch))
def _loss_output(
self,
task_loss: Tensor,
metrics: Dict[str, Tensor],
aux_loss: Optional[Tensor] = 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
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.
@@ -163,17 +193,17 @@ class BaseStrategy(ABC):
if self._rollout_runner is not None:
self._rollout_runner.step()
def __call__(self, batch: Dict[str, Tensor]) -> Tensor:
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(batch)
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(train_batch)
return self.compute_loss_output(train_batch)
class StrategyFactory(BaseFactory["BaseStrategy"]):
@@ -213,9 +243,13 @@ class SEQStrategy(BaseStrategy):
self.label_smoothing = label_smoothing
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
return self.compute_loss_output(batch)["loss"]
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
batch = move_to_device(batch, self.device)
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(
input=logits.flatten(0, 1).float(),
@@ -223,7 +257,7 @@ class SEQStrategy(BaseStrategy):
label_smoothing=self.label_smoothing,
)
return loss
return self._loss_output(loss, {"task_loss": loss}, outputs.get("aux_loss"))
@StrategyFactory.register("sft")
@@ -244,6 +278,9 @@ class SFTStrategy(BaseStrategy):
self.label_smoothing = label_smoothing
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
return self.compute_loss_output(batch)["loss"]
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
batch = move_to_device(batch, self.device)
input_ids, target_ids, position_ids, loss_mask = (
batch["input_ids"],
@@ -255,9 +292,10 @@ class SFTStrategy(BaseStrategy):
ignore_index = -100
input_mask = make_doc_boundary_mask(position_ids)
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
)["logits"]
)
logits = outputs["logits"]
loss = F.cross_entropy(
input=logits.flatten(0, 1).float(),
@@ -266,7 +304,7 @@ class SFTStrategy(BaseStrategy):
label_smoothing=self.label_smoothing,
)
return loss
return self._loss_output(loss, {"task_loss": loss}, outputs.get("aux_loss"))
@StrategyFactory.register("dpo")
@@ -292,6 +330,9 @@ class DPOStrategy(BaseStrategy):
self.reduction = reduction
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
return self.compute_loss_output(batch)["loss"]
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
batch = move_to_device(batch, self.device)
chosen_ids, rejected_ids = batch["chosen"], batch["rejected"]
chosen_mask, rejected_mask = batch["chosen_mask"], batch["rejected_mask"]
@@ -307,22 +348,25 @@ class DPOStrategy(BaseStrategy):
)[None, None, :, :] # [1, 1, S, S]
full_mask = key_pad & causal # [B*2, 1, S, S] — composed
log_pi = get_logprobs(
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():
log_ref = get_logprobs(
ref_output = get_logprobs(
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_rejected = log_pi[chosen_ids.shape[0] :]
@@ -335,7 +379,7 @@ class DPOStrategy(BaseStrategy):
ratio_diff = pi_log_ratio - ref_log_ratio
dpo_loss = -F.logsigmoid(self.beta * ratio_diff).mean()
return dpo_loss
return self._loss_output(dpo_loss, {"dpo_loss": dpo_loss}, aux_loss)
def supports_online(self) -> bool:
return True
@@ -401,9 +445,16 @@ class GRPOStrategy(BaseStrategy):
def sync_old_model(self):
"""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:
return self.compute_loss_output(batch)["loss"]
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
batch = move_to_device(batch, self.device)
prompts = batch["prompts"]
responses = batch["responses"]
@@ -444,16 +495,23 @@ class GRPOStrategy(BaseStrategy):
# get_logprobs returns [B*G, S-1] (S = prompt_len + response_len).
# Response token logprobs occupy the last ``response_len`` positions
# (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, 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():
token_log_probs_old = get_logprobs(
old_output = get_logprobs(
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"]
token_log_probs_old = token_log_probs_old[:, prompt_len - 1 :]
ref_output = get_logprobs(
self.ref_model, full_sequences, attn_mask, full_masks, "none"
)[:, prompt_len - 1 :]
)
token_log_probs_ref = ref_output["logprobs"]
token_log_probs_ref = token_log_probs_ref[:, prompt_len - 1 :]
# Reshape to [B, G, response_len]
token_log_probs_policy = token_log_probs_policy.view(batch_size, group_size, -1)
@@ -486,9 +544,12 @@ class GRPOStrategy(BaseStrategy):
kl_per_token = r - torch.log(r + eps) - 1.0
kl_penalty = self.kl_coef * (kl_per_token * token_masks).sum() / token_count
total_loss = policy_loss + kl_penalty
return total_loss
task_loss = policy_loss + kl_penalty
return self._loss_output(
task_loss,
{"policy_loss": policy_loss, "kl_loss": kl_penalty},
aux_loss,
)
def supports_online(self) -> bool:
return True
@@ -510,5 +571,5 @@ class GRPOStrategy(BaseStrategy):
# 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._entries["online_grpo"] = GRPOStrategy
StrategyFactory._entries["online_dpo"] = DPOStrategy
StrategyFactory.register("online_grpo")(GRPOStrategy)
StrategyFactory.register("online_dpo")(DPOStrategy)
+31 -11
View File
@@ -18,6 +18,7 @@ from astrai.parallel.setup import get_current_device
from astrai.serialization import Checkpoint
from astrai.trainer.metric_util import (
ctx_get_grad_norm,
ctx_get_grad_snr,
ctx_get_loss,
ctx_get_lr,
ctx_get_val_loss,
@@ -235,7 +236,7 @@ class ProgressBarCallback(TrainCallback):
class MetricCallback(TrainCallback):
def __init__(
self,
log_dir: str,
ckpt_dir: str,
save_interval: int,
metrics: List[str] = None,
val_step: int = 0,
@@ -246,8 +247,7 @@ class MetricCallback(TrainCallback):
self.val_step = val_step
self._next_val_step = 0
self.log_dir = Path(log_dir) if log_dir else Path.cwd() / "logs"
self.log_dir.mkdir(parents=True, exist_ok=True)
self.ckpt_dir = Path(ckpt_dir) if ckpt_dir else Path.cwd() / "checkpoint"
self.log_cache = []
@@ -256,14 +256,32 @@ class MetricCallback(TrainCallback):
"lr": ctx_get_lr,
"val_loss": ctx_get_val_loss,
"grad_norm": ctx_get_grad_norm,
"grad_snr": ctx_get_grad_snr,
}
def _metrics(self, context: TrainContext, names):
return {
m: self._metric_funcs[m](context)
for m in names
if self._metric_funcs[m](context) is not None
}
metrics = dict(context.metrics)
for name in names:
metric_fn = self._metric_funcs.get(name)
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)
def _append(self, event_type: str, context: TrainContext, **extra):
@@ -285,8 +303,8 @@ class MetricCallback(TrainCallback):
with torch.no_grad():
for batch in context.val_dataloader:
loss = context.strategy(batch)
total_loss += loss.item()
loss_output = context.strategy(batch)
total_loss += loss_output["loss"].item()
num_batches += 1
if context.world_size > 1 and dist.is_initialized():
@@ -306,13 +324,15 @@ class MetricCallback(TrainCallback):
@only_on_rank(0)
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)
with open(log_file, "w") as f:
for log in self.log_cache:
f.write(json.dumps(log) + "\n")
def on_optimizer_step(self, context):
context.grad_snr_tracker.update(context.model)
if (
context.val_dataloader is not None
and self.val_step > 0
+38 -24
View File
@@ -1,3 +1,4 @@
import logging
import threading
from dataclasses import dataclass, field
from pathlib import Path
@@ -11,13 +12,16 @@ from astrai.config.train_config import TrainConfig
from astrai.dataset import RDSampler
from astrai.inference.core.scheduler import InferenceScheduler
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.protocols import OptimizerProtocol, SchedulerProtocol
from astrai.serialization import Checkpoint, load_json
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, create_ref_model
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
logger = logging.getLogger(__name__)
@dataclass
@@ -34,7 +38,9 @@ class TrainContext:
epoch: int = field(default=0)
consumed_samples: int = field(default=0)
loss: float = field(default=0.0)
metrics: Dict[str, float] = field(default_factory=dict)
grad_norm: Optional[float] = field(default=None)
grad_snr_tracker: GradSNRTracker = field(default_factory=GradSNRTracker)
val_dataloader: Optional[DataLoader] = field(default=None)
val_loss: Optional[float] = field(default=None)
@@ -101,18 +107,13 @@ class TrainContextBuilder:
if checkpoint.config:
model_config = checkpoint.config
if self._resume:
preloaded_epoch = checkpoint.epoch or cfg.start_epoch
if checkpoint.consumed_samples > 0:
per_step = (
cfg.batch_per_device
* get_world_size()
* cfg.grad_accum_steps
)
preloaded_consumed = (
checkpoint.consumed_samples // per_step
) * per_step
else:
preloaded_consumed = cfg.start_samples * get_world_size()
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"):
@@ -131,6 +132,12 @@ class TrainContextBuilder:
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(
world_size=get_world_size(),
rank=get_rank(),
@@ -147,6 +154,7 @@ class TrainContextBuilder:
cfg.optimizer_fn,
cfg.scheduler_fn,
before_wrap=_before_wrap,
after_wrap=_after_wrap,
)
train_dataset = cfg.dataset
@@ -162,6 +170,15 @@ class TrainContextBuilder:
)
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(
data_source=train_dataset,
start_epoch=context.epoch,
@@ -205,6 +222,7 @@ class TrainContextBuilder:
obj.load_state_dict(extra[name])
strategy_kwargs = dict(cfg.extra_kwargs)
strategy_kwargs.setdefault("moe_aux_loss_coef", cfg.moe_aux_loss_coef)
needs_ref = cfg.strategy in (
"dpo",
@@ -215,17 +233,14 @@ class TrainContextBuilder:
needs_old = cfg.strategy in ("grpo", "online_grpo")
if needs_ref:
ref_model = create_ref_model(
cfg.model_fn, executor.unwrap_model(context.model)
).to(device=device)
strategy_kwargs["ref_model"] = ref_model
strategy_kwargs["ref_model"] = create_ref_model(
cfg.model_fn, executor=executor, model=context.model, device=device
)
old_model = None
if needs_old:
old_model = create_ref_model(
cfg.model_fn, executor.unwrap_model(context.model)
).to(device=device)
strategy_kwargs["old_model"] = old_model
strategy_kwargs["old_model"] = create_ref_model(
cfg.model_fn, executor=executor, model=context.model, device=device
)
context.strategy = StrategyFactory.create(
cfg.strategy,
@@ -257,7 +272,6 @@ class TrainContextBuilder:
tokenizer=tokenizer,
max_batch_size=rollout_batch_size,
max_seq_len=max_seq_len,
max_prompt_len=max_seq_len or 4096,
)
generator = RolloutGenerator(
+6 -5
View File
@@ -5,7 +5,7 @@ import torch.distributed as dist
from astrai.config import TrainConfig
from astrai.parallel.setup import spawn_parallel_fn
from astrai.parallel.signal_handler import (
from astrai.signal_handler import (
register_signal_handlers,
unregister_signal_handlers,
)
@@ -42,7 +42,7 @@ class Trainer:
),
CallbackFactory.create(
"metric",
log_dir=cfg.log_dir,
ckpt_dir=cfg.ckpt_dir,
save_interval=cfg.ckpt_interval,
metrics=cfg.metrics,
val_step=cfg.val_step,
@@ -82,9 +82,10 @@ class Trainer:
break
with executor.accumulate(context.model):
self._call_callbacks("on_batch_begin", context)
loss = context.strategy(batch)
context.loss = loss.item()
stand_loss = loss / executor.grad_accum_steps
loss_output = context.strategy(batch)
context.loss = loss_output["loss"].item()
context.metrics = loss_output["metrics"]
stand_loss = loss_output["loss"] / executor.grad_accum_steps
executor.backward(stand_loss)
context.consumed_samples += (
context.config.batch_per_device * context.world_size
+29 -1
View File
@@ -1,6 +1,32 @@
from pathlib import Path
def cuda_toolkit_version() -> tuple[int, int] | None:
"""Return ``(major, minor)`` of the nvcc on PATH, or ``None``.
Used by ``setup.py`` to detect nvcc/torch CUDA version mismatches
(e.g. nvcc 13.0 with a cu128 torch wheel) which cause cryptic ABI errors.
"""
import shutil
import subprocess
nvcc = shutil.which("nvcc")
if nvcc is None:
return None
try:
out = subprocess.check_output(
[nvcc, "--version"], stderr=subprocess.STDOUT, text=True
)
for line in out.splitlines():
if "release" in line:
ver = line.split("release")[1].split(",")[0].strip()
major, minor = ver.split(".")
return (int(major), int(minor))
except Exception:
pass
return None
def _arch_flags() -> list[str]:
import torch
@@ -27,7 +53,7 @@ NVCC_FLAGS = [
"--use_fast_math",
"--ptxas-options=-O3,-v",
"--extra-device-vectorization",
"--threads=8",
"--threads=16",
]
@@ -46,3 +72,5 @@ def register(name: str, sources: list[str] | None = None, **kwargs):
register("attn_decode")
register("attn_prefill")
register("attn_paged_decode")
register("attn_paged_prefill")
register("rotary_emb")
+48 -17
View File
@@ -1,5 +1,13 @@
#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]
};
template<typename T, typename AT = float>
struct AttentionParams {
@@ -19,9 +27,11 @@ struct AttentionParams {
// 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]
int mask_b_stride; // = kv_len (both 2D and 3D)
int mask_q_stride; // 2D: 0 (all q rows share); 3D: kv_len
// Mask: 2D [batch, kv_len], 3D [batch, q_len, kv_len],
// or 4D [batch, n_heads, q_len, kv_len] (head dim broadcasts when stride=0)
int mask_b_stride; // batch stride
int mask_h_stride; // head stride (0 = broadcast across heads)
int mask_q_stride; // q stride (0 = all q rows share)
const T* __restrict__ q;
const T* __restrict__ k;
@@ -33,34 +43,55 @@ struct AttentionParams {
AT* __restrict__ ml_part;
};
// ---- PagedAttentionParams ----
// SGLang-style indirect params over a shared KV pool.
// k_cache/v_cache: [size, kv_head, head_dim] (bare buffers, no gather).
// req_to_token: [num_reqs, max_context_len] token -> slot.
// req_pool_indices:[batch] rows of the current batch into req_to_token.
// kv_indptr: [batch+1] prefix sum of per-request seq_lens (device).
// qo_indptr: [batch+1] prefix sum of per-request q_len (prefill) or
// nullptr for decode (q_len == 1 everywhere).
template<typename T, typename AT = float>
struct PagedAttentionParams {
int batch;
int q_head;
int kv_head;
int q_len;
int kv_len;
int head_dim;
int num_splits;
int use_mask;
int causal_offset;
int causal_offset; // -1 = non-causal; >=0 = causal (per-request offset
// computed inside kernel from kv_indptr/qo_indptr)
float scale;
int num_splits;
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;
// Q: [total_q, q_head, head_dim] (3D flattened — no batch dim).
// For decode total_q == batch (q_len=1 per request).
// For prefill total_q == qo_indptr[batch].
int q_stride_l, q_stride_h, q_stride_d;
// Q: [total_q, q_head, head_dim]
const T* __restrict__ q;
// Flat KV pool: [size, kv_head, head_dim]
const T* __restrict__ k_cache;
const T* __restrict__ v_cache;
// Indexing
const int64_t* __restrict__ req_to_token; // [num_reqs, max_context_len]
const int64_t* __restrict__ req_pool_indices; // [batch]
const int* __restrict__ kv_indptr; // [batch+1]
const int* __restrict__ qo_indptr; // [batch+1] or nullptr (decode)
int max_context_len; // req_to_token stride (dim 1)
int max_seq_len; // max per-request seq_len (host-side, for split computation)
int total_q; // total Q tokens across all requests (host-side, for grid)
int max_q_len; // max per-request q_len (host-side, for prefill grid)
// Mask: [batch, max_seq_len] (decode) or [batch, 1, q_len, kv_len]
// (prefill, optional). mask_h_stride/mask_q_stride are 0 when those
// dims are size 1 (broadcast).
int mask_b_stride;
int mask_h_stride;
int mask_q_stride;
const bool* __restrict__ mask;
const int64_t* __restrict__ page_table;
T* __restrict__ o;
AT* __restrict__ o_part;
+7 -3
View File
@@ -10,17 +10,21 @@ torch::Tensor attn_decode(
double scale,
int64_t layout
) {
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
auto stream = at::cuda::getCurrentCUDAStream();
AttentionParams<bf16> 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.head_dim % 32 == 0, "head_dim must be multiple of 32");
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();
alloc_split_partials(p);
DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p);
DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p, stream);
C10_CUDA_CHECK(cudaGetLastError());
return O;
}
@@ -32,6 +36,6 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
py::arg("mask") = py::none(),
py::arg("causal_offset") = -1,
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)");
}
+5 -5
View File
@@ -24,7 +24,7 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
// KV: [batch, kv_head, kv_len, head_dim] — stride-based base
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
int mask_base = batch * p.mask_b_stride;
int mask_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
@@ -70,8 +70,8 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
}
float new_m = fmaxf(m, partial);
float alpha = expf(m - new_m);
float beta = expf(partial - new_m);
float alpha = __expf(m - new_m);
float beta = __expf(partial - new_m);
d = d * alpha + beta;
int v_off = kv_base + kv_idx * p.kv_stride_l
@@ -116,8 +116,8 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
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);
float corr = __expf(m - nm);
float e = __expf(mi - nm);
acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
l = fmaf(l, corr, li * e);
m = nm;
+31 -23
View File
@@ -76,26 +76,17 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
cp_async_commit();
};
constexpr int BUF_MASK = (Traits::STAGES > 1) ? (Traits::STAGES - 1) : 0;
// Prologue
if (ti_begin < ti_end) {
load_tile(ti_begin, 0);
}
for (int ti = ti_begin; ti < ti_end; ti++) {
int buf = (ti - ti_begin) & BUF_MASK;
cp_async_wait_group<0>();
__syncwarp();
if constexpr (Traits::STAGES > 1) {
if (ti + 1 < ti_end)
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
}
// ---- Multi-stage cp.async pipeline ----
// Prologue loads STAGES tiles; each loop iteration waits only for the
// 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;
auto process_tile = [&](int it, int buf) {
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
int kv0 = ti * Traits::BC;
int kv0 = (ti_begin + it) * Traits::BC;
float Sacc[Traits::NC8][4];
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
@@ -109,18 +100,35 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
0, 0,
p.mask_b_stride, 0,
batch,
p.mask_b_stride, 0, 0,
batch, 0,
p.mask,
Sacc, Oacc, m0, m1, l0, l1, lane);
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
__syncwarp();
};
if constexpr (Traits::STAGES == 1) {
if (ti + 1 < ti_end)
load_tile(ti + 1, 0);
if (ntiles >= STAGES) {
#pragma unroll
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 ----
+115 -87
View File
@@ -8,65 +8,82 @@
#include "attn_prefill_split_q.cuh"
#include "attn_decode_split_kv.cuh"
#include "attn_paged_decode_split_kv.cuh"
#include "attn_paged_prefill_split_q.cuh"
#ifndef ASTRAI_NO_MMA
#include "attn_prefill_split_q_mma.cuh"
#include "attn_decode_split_kv_mma.cuh"
#include "attn_paged_decode_split_kv_mma.cuh"
#include "attn_paged_prefill_split_q_mma.cuh"
#endif
// Split-KV: compute number of splits to fill all SMs for small-batch decode.
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, MAX_SPLITS)));
// 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, launch_decode_mma, HEAD_DIM, p, group_size);
#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
// ======================================================================
#ifndef ASTRAI_NO_MMA
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static inline void launch_prefill_mma(AttentionParams<bf16>& p) {
static inline void launch_prefill_mma(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>;
dim3 grid((p.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, IsCausal, HasMask><<<grid, block>>>(p);
attn_prefill_split_q_mma_kernel<Traits, IsCausal, HasMask><<<grid, block, 0, stream>>>(p);
}
#endif
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static inline void launch_prefill_scalar(AttentionParams<bf16>& p) {
static inline void launch_prefill_scalar(AttentionParams<bf16>& p, cudaStream_t stream) {
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);
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask><<<grid, block>>>(p);
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask><<<grid, block, 0, stream>>>(p);
}
template <int HEAD_DIM>
static inline void dispatch_prefill(AttentionParams<bf16>& p) {
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
if (is_causal) {
if (has_mask) launch_prefill_mma<HEAD_DIM, true, true>(p);
else launch_prefill_mma<HEAD_DIM, true, false>(p);
} else {
if (has_mask) launch_prefill_mma<HEAD_DIM, false, true>(p);
else launch_prefill_mma<HEAD_DIM, false, false>(p);
}
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_prefill_mma, HEAD_DIM, p, stream);
#else
if (is_causal) {
if (has_mask) launch_prefill_scalar<HEAD_DIM, true, true>(p);
else launch_prefill_scalar<HEAD_DIM, true, false>(p);
} else {
if (has_mask) launch_prefill_scalar<HEAD_DIM, false, true>(p);
else launch_prefill_scalar<HEAD_DIM, false, false>(p);
}
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_prefill_scalar, HEAD_DIM, p, stream);
#endif
}
@@ -75,121 +92,132 @@ static inline void dispatch_prefill(AttentionParams<bf16>& p) {
// ======================================================================
#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 <int HEAD_DIM, bool IsCausal, bool HasMask>
static inline void launch_decode_mma(AttentionParams<bf16>& p, int group_size) {
static inline void launch_decode_mma(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;
int tiles_total = (p.kv_len + 32 - 1) / 32;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
using Traits = KernelTraits<HEAD_DIM, 32, 1, STAGES>;
constexpr int BC = 16;
int tiles_total = (p.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, IsCausal, HasMask><<<grid, 32>>>(p);
attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask><<<grid, 32, 0, stream>>>(p);
}
#endif
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static inline void launch_decode_scalar(AttentionParams<bf16>& p, int group_size) {
static inline void launch_decode_scalar(AttentionParams<bf16>& p, int group_size, cudaStream_t stream) {
int chunks_total = (p.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, IsCausal, HasMask><<<grid, block, smem>>>(p);
attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem, stream>>>(p);
}
template <int HEAD_DIM>
static inline void dispatch_decode(AttentionParams<bf16>& p) {
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
if (is_causal) {
if (has_mask) launch_decode_mma<HEAD_DIM, true, true>(p, group_size);
else launch_decode_mma<HEAD_DIM, true, false>(p, group_size);
} else {
if (has_mask) launch_decode_mma<HEAD_DIM, false, true>(p, group_size);
else launch_decode_mma<HEAD_DIM, false, false>(p, group_size);
}
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_mma, HEAD_DIM, p, group_size, stream);
#else
if (is_causal) {
if (has_mask) launch_decode_scalar<HEAD_DIM, true, true>(p, group_size);
else launch_decode_scalar<HEAD_DIM, true, false>(p, group_size);
} else {
if (has_mask) launch_decode_scalar<HEAD_DIM, false, true>(p, group_size);
else launch_decode_scalar<HEAD_DIM, false, false>(p, group_size);
}
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_scalar, HEAD_DIM, p, group_size, stream);
#endif
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
}
// ======================================================================
// Paged Decode
// Paged Decode (SGLang-style: flat pool + req_to_token + kv_indptr)
// ======================================================================
#ifndef ASTRAI_NO_MMA
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static inline void launch_paged_decode_mma(PagedAttentionParams<bf16>& p, int group_size) {
static inline void launch_paged_decode_mma(PagedAttentionParams<bf16>& p, cudaStream_t stream) {
int G = p.q_head / p.kv_head;
constexpr int MAX_G = 16;
bool page_ok = (p.page_size >= 32);
if (G >= 1 && page_ok) {
int num_passes = (G + MAX_G - 1) / MAX_G;
int tiles_total = (p.kv_len + 32 - 1) / 32;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
using Traits = KernelTraits<HEAD_DIM, 32, 1, STAGES>;
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
paged_attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask> <<<grid, 32>>>(p);
} else {
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);
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
dim3 block(32, group_size);
paged_attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
}
constexpr int BC = 16;
int num_passes = (G + MAX_G - 1) / MAX_G;
int tiles_total = (p.max_seq_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);
paged_attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask> <<<grid, 32, 0, stream>>>(p);
}
#endif
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static inline void launch_paged_decode_scalar(PagedAttentionParams<bf16>& p, int group_size) {
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
static inline void launch_paged_decode_scalar(PagedAttentionParams<bf16>& p, int group_size, cudaStream_t stream) {
int chunks_total = (p.max_seq_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);
int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit
int g = min(group_size, 32);
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
dim3 block(32, g);
paged_attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
paged_attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem, stream>>>(p);
}
template <int HEAD_DIM>
static inline void dispatch_paged_decode(PagedAttentionParams<bf16>& p) {
static inline void dispatch_paged_decode(PagedAttentionParams<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
if (is_causal) {
if (has_mask) launch_paged_decode_mma<HEAD_DIM, true, true>(p, group_size);
else launch_paged_decode_mma<HEAD_DIM, true, false>(p, group_size);
} else {
if (has_mask) launch_paged_decode_mma<HEAD_DIM, false, true>(p, group_size);
else launch_paged_decode_mma<HEAD_DIM, false, false>(p, group_size);
}
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_decode_mma, HEAD_DIM, p, stream);
#else
if (is_causal) {
if (has_mask) launch_paged_decode_scalar<HEAD_DIM, true, true>(p, group_size);
else launch_paged_decode_scalar<HEAD_DIM, true, false>(p, group_size);
} else {
if (has_mask) launch_paged_decode_scalar<HEAD_DIM, false, true>(p, group_size);
else launch_paged_decode_scalar<HEAD_DIM, false, false>(p, group_size);
}
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_decode_scalar, HEAD_DIM, p, group_size, stream);
#endif
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
}
// ======================================================================
// Paged Prefill (SGLang-style: flat pool + ragged batch)
// ======================================================================
#ifndef ASTRAI_NO_MMA
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static inline void launch_paged_prefill_mma(PagedAttentionParams<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 max_q_tiles = (p.max_q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS);
dim3 grid(max_q_tiles, p.q_head, p.batch);
dim3 block(Traits::NUM_THREADS);
paged_attn_prefill_split_q_mma_kernel<Traits, IsCausal, HasMask><<<grid, block, 0, stream>>>(p);
}
#endif
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static inline void launch_paged_prefill_scalar(PagedAttentionParams<bf16>& p, cudaStream_t stream) {
constexpr int G = 8, ROWS = 32, P_BC = 32;
int max_q_tiles = (p.max_q_len + ROWS - 1) / ROWS;
dim3 grid(max_q_tiles, p.q_head, p.batch);
dim3 block(G, ROWS);
paged_attn_prefill_split_q_kernel<HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask>
<<<grid, block, 0, stream>>>(p);
}
template <int HEAD_DIM>
static inline void dispatch_paged_prefill(PagedAttentionParams<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, launch_paged_prefill_mma, HEAD_DIM, p, stream);
#else
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_prefill_scalar, HEAD_DIM, p, stream);
#endif
}
+177 -37
View File
@@ -1,4 +1,5 @@
#pragma once
#include <float.h>
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include "attn_common.h"
@@ -7,24 +8,28 @@
using bf16 = __nv_bfloat16;
// Dispatch head_dim: shared macro — avoids C++20 lambda template syntax.
// Usage: DISPATCH_HEAD_DIM(hd, fn, arg)
// Expands to: fn<32>(arg); fn<64>(arg); etc.
#define DISPATCH_HEAD_DIM(hd, fn, arg) \
// Usage: DISPATCH_HEAD_DIM(hd, fn, args...)
// Expands to: fn<32>(args...); fn<64>(args...); etc.
#define DISPATCH_HEAD_DIM(hd, fn, ...) \
switch (hd) { \
case 32: fn<32>(arg); break; \
case 64: fn<64>(arg); break; \
case 128: fn<128>(arg); break; \
case 256: fn<256>(arg); break; \
case 32: fn<32>(__VA_ARGS__); break; \
case 64: fn<64>(__VA_ARGS__); break; \
case 128: fn<128>(__VA_ARGS__); break; \
case 256: fn<256>(__VA_ARGS__); break; \
default: \
TORCH_CHECK(false, "unsupported head_dim ", hd, \
" (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>
inline void alloc_split_partials(P& p) {
auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
auto o_part = torch::empty({p.batch, p.q_head, MAX_SPLITS, p.head_dim}, fopt);
auto ml_part = torch::empty({p.batch, p.q_head, MAX_SPLITS, 2}, fopt);
auto o_part = torch::empty(at::IntArrayRef{p.batch, p.q_head, MAX_SPLITS, p.head_dim}, 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.ml_part = (float*)ml_part.data_ptr();
}
@@ -32,7 +37,7 @@ inline void alloc_split_partials(P& p) {
// ---- Shared Q-dims + strides extraction ----
template <typename 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.q_head = (int)q.size(1);
p.q_len = (int)q.size(2);
@@ -44,6 +49,9 @@ inline void extract_q_dims_and_strides(torch::Tensor& q, int64_t layout, P& p) {
}
// ---- 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>
inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
if (p.use_mask) {
@@ -54,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");
if (m.dim() == 2) {
p.mask_b_stride = (int)m.stride(0);
p.mask_h_stride = 0;
p.mask_q_stride = 0;
} 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_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 {
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>();
} else {
p.mask = nullptr;
p.mask_b_stride = 0;
p.mask_h_stride = 0;
p.mask_q_stride = 0;
}
}
@@ -93,7 +109,7 @@ inline void attn_pack_params(
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_len = (int)k.size(2);
@@ -118,54 +134,178 @@ inline void attn_pack_params(
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>
inline void attn_pack_paged_params(
inline void attn_pack_paged_decode_params(
torch::Tensor q,
torch::Tensor page_table,
torch::Tensor k_cache,
torch::Tensor v_cache,
int64_t page_size,
int64_t kv_len,
torch::Tensor req_to_token,
torch::Tensor req_pool_indices,
torch::Tensor kv_indptr,
int64_t max_seq_len,
c10::optional<torch::Tensor> mask,
int64_t causal_offset,
double scale,
int64_t layout,
PagedAttentionParams<T>& p
) {
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(k_cache.dtype() == torch::kBFloat16, "k_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(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must have identical shapes");
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(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.kv_head = (int)k_cache.size(2);
p.kv_len = (int)kv_len;
p.page_size = (int)page_size;
p.max_pages = (int)page_table.size(1);
TORCH_CHECK(q.size(2) == 1, "Q seq_len must be 1 (decode)");
p.batch = (int)q.size(0);
p.q_head = (int)q.size(1);
p.head_dim = (int)q.size(2);
p.kv_head = (int)k_cache.size(1);
TORCH_CHECK(k_cache.size(2) == p.head_dim, "k_cache head_dim mismatch");
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
TORCH_CHECK(k_cache.size(1) == page_size,
"k_cache dim 1 must equal page_size, got ",
k_cache.size(1), " vs ", page_size);
TORCH_CHECK(p.q_head % p.kv_head == 0, "q_head must be divisible by kv_head");
p.q_stride_l = (int)q.stride(0);
p.q_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.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.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,
PagedAttentionParams<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.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 = 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_part = nullptr;
p.ml_part = nullptr;
pack_mask(mask, p);
}
+10 -4
View File
@@ -3,6 +3,12 @@
#include <cuda_fp16.h>
#include <cuda_runtime.h>
// Predicated cp.async (4-operand form) requires CUDA 11.2+.
// 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.
//
@@ -192,8 +198,8 @@ __device__ inline void mma_softmax_tile(
int kv0,
int maxc0, int maxc1,
int qrow0, int qrow1,
int mask_b_stride, int mask_q_stride,
int mask_batch,
int mask_b_stride, int mask_h_stride, int mask_q_stride,
int mask_batch, int mask_head,
const bool* __restrict__ mask,
float Sacc[Traits::NC8][4],
float Oacc[Traits::DN8][4],
@@ -204,8 +210,8 @@ __device__ inline void mma_softmax_tile(
int tid4 = lane & 3;
float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX;
int mask_base0 = mask_batch * mask_b_stride + qrow0 * mask_q_stride;
int mask_base1 = mask_batch * mask_b_stride + qrow1 * 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 + mask_head * mask_h_stride + qrow1 * mask_q_stride;
#pragma unroll
for (int n8 = 0; n8 < Traits::NC8; n8++) {
int cc = kv0 + n8 * 8 + 2 * tid4;
+21 -17
View File
@@ -3,40 +3,44 @@
torch::Tensor attn_paged_decode(
torch::Tensor q,
torch::Tensor page_table,
torch::Tensor k_cache,
torch::Tensor v_cache,
int64_t page_size,
int64_t kv_len,
torch::Tensor req_to_token,
torch::Tensor req_pool_indices,
torch::Tensor kv_indptr,
int64_t max_seq_len,
c10::optional<torch::Tensor> mask,
int64_t causal_offset,
double scale,
int64_t layout
double scale
) {
PagedAttentionParams<bf16> p;
attn_pack_paged_params(q, page_table, k_cache, v_cache,
page_size, kv_len, mask, causal_offset, scale, layout, p);
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
auto stream = at::cuda::getCurrentCUDAStream();
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
p.o = (bf16*)O_view.data_ptr();
PagedAttentionParams<bf16> p;
attn_pack_paged_decode_params(q, k_cache, v_cache,
req_to_token, req_pool_indices, kv_indptr,
max_seq_len, mask, 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();
alloc_split_partials(p);
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p);
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p, stream);
C10_CUDA_CHECK(cudaGetLastError());
return O;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("attn_paged_decode", &attn_paged_decode,
py::arg("q"),
py::arg("page_table"),
py::arg("k_cache"),
py::arg("v_cache"),
py::arg("page_size"),
py::arg("kv_len"),
py::arg("req_to_token"),
py::arg("req_pool_indices"),
py::arg("kv_indptr"),
py::arg("max_seq_len"),
py::arg("mask") = py::none(),
py::arg("causal_offset") = -1,
py::arg("scale") = 0.0,
py::arg("layout") = 0,
"Paged GQA decode — split-KV with direct page-table access.");
"SGLang-style paged decode: flat KV pool + req_to_token + kv_indptr.");
}
+33 -28
View File
@@ -5,6 +5,8 @@
#include "attn_warp_utils.cuh"
constexpr int PDC_CHUNK = 64;
// Scalar paged decode (fallback for sm < 80, no tensor cores).
// Reads K/V from flat pool via req_to_token indexing.
template <int HEAD_DIM, bool IsCausal, bool HasMask>
__global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p) {
int batch = blockIdx.x / p.kv_head;
@@ -15,8 +17,11 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
int lane = threadIdx.x;
int hd_per_thread = p.head_dim / 32;
const int seq_len = p.kv_indptr[batch + 1] - p.kv_indptr[batch];
const int64_t req_idx = p.req_pool_indices[batch];
float q_reg[8];
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
int q_off = batch * p.q_stride_l + 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++)
@@ -26,16 +31,19 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
extern __shared__ __align__(16) bf16 k_smem[];
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
int chunks_total = (seq_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;
const int64_t pool_stride = (int64_t)p.kv_head * p.head_dim;
const int64_t head_off = (int64_t)kv_head * p.head_dim;
const int64_t rtt_stride = (int64_t)p.max_context_len;
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 this_chunk = min(PDC_CHUNK, seq_len - chunk_start);
int total = this_chunk * p.head_dim;
for (int i = threadIdx.y * 32 + lane; i < total;
@@ -43,14 +51,9 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
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;
int64_t slot = p.req_to_token[req_idx * rtt_stride + pos];
if (slot >= 0) {
int64_t off = slot * pool_stride + head_off + d_dim;
k_smem[i] = p.k_cache[off];
} else {
k_smem[i] = __float2bfloat16(0.0f);
@@ -67,28 +70,30 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
partial = warp_reduce_sum(partial) * p.scale;
int kv_idx = chunk_start + s;
bool masked = false;
if constexpr (HasMask) {
if (!p.mask[mask_base + kv_idx])
partial = -FLT_MAX;
}
if constexpr (IsCausal) {
if (kv_idx > p.causal_offset)
partial = -FLT_MAX;
masked = true;
}
// Decode: the query is the last token, so its valid range [0,
// seq_len) IS the causal range. IsCausal is accepted for dispatch
// uniformity but must not apply causal_offset masking here.
if (masked)
partial = -FLT_MAX;
float new_m = fmaxf(m, partial);
float alpha = expf(m - new_m);
float beta = expf(partial - new_m);
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;
int64_t slot = p.req_to_token[req_idx * rtt_stride + pos];
if (masked) {
#pragma unroll
for (int i = 0; i < hd_per_thread; i++)
acc_reg[i] = fmaf(acc_reg[i], alpha, 0.0f);
} else if (slot >= 0) {
int64_t v_base = slot * pool_stride + head_off;
#pragma unroll
for (int i = 0; i < hd_per_thread; i++)
acc_reg[i] = fmaf(acc_reg[i], alpha,
@@ -133,14 +138,14 @@ __global__ void paged_attn_decode_combine_kernel(PagedAttentionParams<bf16> p) {
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);
float corr = __expf(m - nm);
float e = __expf(mi - nm);
acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
l = fmaf(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;
int o_off = batch * p.q_stride_l + q_head * p.q_stride_h + d * p.q_stride_d;
p.o[o_off] = __float2bfloat16(acc * inv);
}
+61 -44
View File
@@ -5,12 +5,16 @@
#include "attn_mma_utils.cuh"
#include "attn_warp_utils.cuh"
// Paged split-KV tensor-core decode via GQA head-packing.
// Reads K/V directly from the page pool through a page table — one tile
// (BC=32) fits within a single page (page_size >= 32), so the page-table
// lookup happens once per tile for cp.async.
// SGLang-style split-KV tensor-core decode.
//
// IsCausal and HasMask are compile-time bools.
// Reads K/V directly from a flat pool [size, kv_head, head_dim] via
// req_to_token indexing — no gather, no page-table dimension.
// Each batch element has its own seq_len (from kv_indptr), eliminating
// padding waste: short sequences only process the tiles they own.
//
// For decode (q_len=1), causal masking is implicit — each request attends
// to [0, seq_len) which is exactly its valid range. The IsCausal flag
// is accepted for dispatch uniformity but does not change maxc.
template <typename Traits, bool IsCausal, bool HasMask>
__global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16> p) {
const int lane = threadIdx.x;
@@ -22,6 +26,10 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
const int batch = blockIdx.y;
const int split = blockIdx.z;
// Per-request seq_len from device-side kv_indptr — no padding.
const int seq_len = p.kv_indptr[batch + 1] - p.kv_indptr[batch];
const int64_t req_idx = p.req_pool_indices[batch];
constexpr int MAX_G = 16;
const int G_total = p.q_head / p.kv_head;
const int g_begin = pass * MAX_G;
@@ -31,13 +39,14 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
const int q_base = batch * p.q_stride_l + 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[Traits::KD][4];
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
qra, qrb, va, vb, tid4, Qa);
load_q_mma_frags<Traits::KD>(p.q + q_base,
p.q_stride_h, p.q_stride_d,
qra, qrb, va, vb, tid4, Qa);
float Oacc[Traits::DN8][4];
#pragma unroll
@@ -45,33 +54,35 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
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 + Traits::BC - 1) / Traits::BC;
const int tiles_total = (seq_len + Traits::BC - 1) / Traits::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 int64_t page_stride = (int64_t)p.page_size * p.kv_head * Traits::HEAD_DIM;
const int64_t pos_stride = (int64_t)p.kv_head * Traits::HEAD_DIM;
// Flat pool stride: [size, kv_head, head_dim] — contiguous.
const int64_t pool_stride = (int64_t)p.kv_head * Traits::HEAD_DIM;
const int64_t head_off = (int64_t)kv_head * Traits::HEAD_DIM;
const int64_t rtt_stride = (int64_t)p.max_context_len;
// ---- Load tile lambda: paged addressing ----
// ---- Load tile lambda: SGLang addressing ----
// slot = req_to_token[req_idx * max_context_len + kc]
// gmem = k_cache[slot * pool_stride + head_off + d]
auto load_tile = [&](int ti, int buf) {
int kv0 = ti * Traits::BC;
bf16* dK = sK + buf * Traits::BC * Traits::LD;
bf16* dV = sV + buf * Traits::BC * Traits::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 * Traits::VEC; i < Traits::TOTAL;
i += Traits::NUM_THREADS * Traits::VEC) {
int r = i / Traits::HEAD_DIM, d = i % Traits::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;
bool valid = (kc < seq_len);
if constexpr (HasMask) {
valid = valid && p.mask[batch * p.mask_b_stride + kc];
}
int64_t slot = valid ? p.req_to_token[req_idx * rtt_stride + kc] : 0;
valid = valid && (slot >= 0);
int64_t gmem_base = slot * pool_stride + head_off;
int off = r * Traits::LD + swiz_col(d, r, Traits::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);
@@ -79,25 +90,13 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
cp_async_commit();
};
constexpr int BUF_MASK = (Traits::STAGES > 1) ? (Traits::STAGES - 1) : 0;
if (ti_begin < ti_end) {
load_tile(ti_begin, 0);
}
for (int ti = ti_begin; ti < ti_end; ti++) {
int buf = (ti - ti_begin) & BUF_MASK;
cp_async_wait_group<0>();
__syncwarp();
if constexpr (Traits::STAGES > 1) {
if (ti + 1 < ti_end)
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
}
constexpr int STAGES = Traits::STAGES;
const int ntiles = ti_end - ti_begin;
auto process_tile = [&](int it, int buf) {
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
int kv0 = ti * Traits::BC;
int kv0 = (ti_begin + it) * Traits::BC;
float Sacc[Traits::NC8][4];
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
@@ -107,23 +106,41 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
// For decode, maxc = seq_len regardless of IsCausal — the valid
// range [0, seq_len) IS the causal range (query is the last token).
mma_softmax_tile<Traits, HasMask>(kv0, seq_len, seq_len,
0, 0,
p.mask_b_stride, 0,
batch,
p.mask_b_stride, 0, 0,
batch, 0,
p.mask,
Sacc, Oacc, m0, m1, l0, l1, lane);
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
__syncwarp();
};
if constexpr (Traits::STAGES == 1) {
if (ti + 1 < ti_end)
load_tile(ti + 1, 0);
if (ntiles >= STAGES) {
#pragma unroll
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 {
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 partials ----
auto split_slot = [&](int h) -> size_t {
size_t bh = (size_t)batch * p.q_head + h;
return bh * MAX_SPLITS + split;
+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();
PagedAttentionParams<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.");
}
+126
View File
@@ -0,0 +1,126 @@
#pragma once
#include <cuda_bf16.h>
#include <float.h>
#include "attn_common.h"
using bf16 = __nv_bfloat16;
// Scalar paged prefill (fallback for sm < 80, no tensor cores).
// Reads K/V from a flat pool via req_to_token, supports ragged batches
// via qo_indptr + kv_indptr. Mirrors the split-Q MMA kernel's indexing:
// grid (max_q_tiles, q_head, batch), block (G, ROWS).
//
// HasMask: 4D mask [batch, 1, q_len, kv_len] (True=keep), columns are
// request-local kv positions. q_head is the q-index (mask_h broadcast).
//
// group_reduce_sum<G> is provided by attn_prefill_split_q.cuh (already
// included via the dispatcher).
template <int HEAD_DIM, int G, int ROWS, int P_BC, bool IsCausal, bool HasMask>
__global__ void paged_attn_prefill_split_q_kernel(PagedAttentionParams<bf16> p) {
constexpr int DPT = HEAD_DIM / G;
const int q_tile = blockIdx.x;
const int q_head = blockIdx.y;
const int req_b = blockIdx.z;
const int gpos = threadIdx.x; // 0..G-1 (d-chunk)
const int row = threadIdx.y; // 0..ROWS-1 (q row within tile)
const int q_row = q_tile * ROWS + row;
const int seq_len = p.kv_indptr[req_b + 1] - p.kv_indptr[req_b];
const int q_len = p.qo_indptr[req_b + 1] - p.qo_indptr[req_b];
const int causal_off = seq_len - q_len;
const int64_t req_idx = p.req_pool_indices[req_b];
const int kv_head = q_head / (p.q_head / p.kv_head);
__shared__ __align__(16) bf16 sK[P_BC * HEAD_DIM];
__shared__ __align__(16) bf16 sV[P_BC * HEAD_DIM];
// Q base: absolute token = qo_indptr[req_b] + q_row.
float qreg[DPT];
if (q_row < q_len) {
int q_off = (p.qo_indptr[req_b] + q_row) * p.q_stride_l
+ q_head * p.q_stride_h + gpos * DPT * p.q_stride_d;
#pragma unroll
for (int i = 0; i < DPT; i++)
qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
}
float m = -FLT_MAX, l = 0.0f, acc[DPT];
#pragma unroll
for (int i = 0; i < DPT; i++) acc[i] = 0.0f;
const int64_t pool_stride = (int64_t)p.kv_head * p.head_dim;
const int64_t head_off = (int64_t)kv_head * p.head_dim;
const int64_t rtt_stride = (int64_t)p.max_context_len;
const int mask_base = req_b * p.mask_b_stride + q_head * p.mask_h_stride
+ q_row * p.mask_q_stride;
int tiles = (seq_len + P_BC - 1) / P_BC;
int tt = G * ROWS;
int lid = row * G + gpos;
// Each warp holds (32/G) q-rows; reduce only within this row's G lanes.
int lane_in_warp = lid & 31;
unsigned gmask = (G == 32) ? 0xFFFFFFFFu
: (((1u << G) - 1u) << (lane_in_warp & ~(G - 1)));
for (int ti = 0; ti < tiles; ti++) {
int kv0 = ti * P_BC;
int tlen = min(P_BC, seq_len - kv0);
// Load K/V tile into shared memory via req_to_token (request-local pos).
for (int i = lid; i < tlen * HEAD_DIM; i += tt) {
int s = i / HEAD_DIM, d_dim = i % HEAD_DIM;
int pos = kv0 + s;
int64_t slot = p.req_to_token[req_idx * rtt_stride + pos];
int64_t off = slot * pool_stride + head_off + d_dim;
sK[i] = (slot >= 0) ? p.k_cache[off] : __float2bfloat16(0.0f);
sV[i] = (slot >= 0) ? p.v_cache[off] : __float2bfloat16(0.0f);
}
__syncthreads();
int lim = tlen;
if constexpr (IsCausal) {
if (q_row < q_len) {
int ep = causal_off + q_row + 1;
if (kv0 >= ep)
lim = 0;
else if (kv0 + tlen > ep)
lim = ep - kv0;
}
}
for (int s = 0; s < lim; s++) {
bool keep = true;
if constexpr (HasMask) {
if (q_row < q_len && !p.mask[mask_base + kv0 + s])
keep = false;
}
float w = 0.0f;
#pragma unroll
for (int i = 0; i < DPT; i++)
w += qreg[i] * __bfloat162float(sK[s * HEAD_DIM + gpos * DPT + i]);
w = group_reduce_sum<G>(w, gmask) * p.scale;
if (!keep) w = -FLT_MAX;
float nm = fmaxf(m, w);
float alpha = __expf(m - nm);
float beta = __expf(w - nm);
l = l * alpha + beta;
#pragma unroll
for (int i = 0; i < DPT; i++)
acc[i] = acc[i] * alpha
+ __bfloat162float(sV[s * HEAD_DIM + gpos * DPT + i]) * beta;
m = nm;
}
__syncthreads();
}
if (q_row >= q_len) return;
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
int o_off = (p.qo_indptr[req_b] + q_row) * p.q_stride_l
+ q_head * p.q_stride_h + gpos * DPT * p.q_stride_d;
#pragma unroll
for (int i = 0; i < DPT; i++)
p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * inv);
}
@@ -0,0 +1,164 @@
#pragma once
#include <cfloat>
#include <cuda_bf16.h>
#include "attn_common.h"
#include "attn_mma_utils.cuh"
// SGLang-style split-Q tensor-core prefill.
//
// Reads K/V directly from a flat pool [size, kv_head, head_dim] via
// req_to_token — no gather, no temporary tensor. Supports ragged batches:
// each request has its own q_len and kv_len, addressed via qo_indptr and
// kv_indptr.
//
// Grid: (max_q_tiles, q_head, batch) — one batch element per blockIdx.z.
// Blocks beyond a request's q_len exit early after writing sentinel-free
// no-ops. This avoids the binary-search approach and guarantees every Q
// token is covered, even when q_len < BR*WARPS (e.g. decode-like prefill).
//
// Q layout: [total_q, q_head, head_dim] (3D, flattened across requests).
// O layout: same as Q.
//
// IsCausal is a compile-time bool. When true, each Q row qi (within its
// request) attends to [0, causal_offset_b + qi + 1) where
// causal_offset_b = kv_len_b - q_len_b (position of first Q token).
template <typename Traits, bool IsCausal, bool HasMask>
__global__ void paged_attn_prefill_split_q_mma_kernel(PagedAttentionParams<bf16> p) {
const int warp = threadIdx.x / 32;
const int lane = threadIdx.x % 32;
const int gid = lane >> 2;
const int tid4 = lane & 3;
const int q_head = blockIdx.y;
const int req_b = blockIdx.z;
const int qrow0 = (blockIdx.x * Traits::WARPS + warp) * Traits::BR;
const int seq_len = p.kv_indptr[req_b + 1] - p.kv_indptr[req_b];
const int q_len = p.qo_indptr[req_b + 1] - p.qo_indptr[req_b];
const int causal_off = seq_len - q_len;
const int64_t req_idx = p.req_pool_indices[req_b];
// No per-warp early exit — all warps must participate in __syncthreads.
// Warps beyond q_len get zero-filled Q frags (va=vb=false) and skip output.
const int kv_head = q_head / (p.q_head / p.kv_head);
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
// Q base: offset by qo_indptr[req_b] to get absolute token address.
const int q_base = p.qo_indptr[req_b] * p.q_stride_l + q_head * p.q_stride_h;
const int qra = qrow0 + gid;
const int qrb = qrow0 + gid + 8;
const bool va = qra < q_len, vb = qrb < q_len;
unsigned Qa[Traits::KD][4];
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_l, p.q_stride_d,
qra, qrb, va, vb, tid4, Qa);
float Oacc[Traits::DN8][4];
#pragma unroll
for (int j = 0; j < Traits::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 int64_t pool_stride = (int64_t)p.kv_head * Traits::HEAD_DIM;
const int64_t head_off = (int64_t)kv_head * Traits::HEAD_DIM;
const int64_t rtt_stride = (int64_t)p.max_context_len;
const int tiles = (seq_len + Traits::BC - 1) / Traits::BC;
const int qr0 = qrow0 + gid;
const int qr1 = qrow0 + gid + 8;
// Causal tile-skip (dead code when IsCausal == false)
const int max_kv = qrow0 + Traits::BR - 1 + causal_off;
const int block_max_kv =
blockIdx.x * Traits::WARPS * Traits::BR + Traits::WARPS * Traits::BR - 1
+ causal_off;
int t_end = tiles - 1;
if constexpr (IsCausal) {
int bt = block_max_kv / Traits::BC;
if (bt < t_end) t_end = bt;
}
// ---- Load tile lambda: SGLang addressing ----
auto load_tile = [&](int ti, int buf) {
int kv0 = ti * Traits::BC;
bf16* dK = sK + buf * Traits::BC * Traits::LD;
bf16* dV = sV + buf * Traits::BC * Traits::LD;
#pragma unroll
for (int i = threadIdx.x * Traits::VEC; i < Traits::TOTAL;
i += Traits::NUM_THREADS * Traits::VEC) {
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
int kc = kv0 + r;
bool valid = kc < seq_len;
int64_t slot = valid ? p.req_to_token[req_idx * rtt_stride + kc] : 0;
valid = valid && (slot >= 0);
int64_t gmem_base = slot * pool_stride + head_off;
int off = r * Traits::LD + swiz_col(d, r, Traits::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 + main loop (FA2-style double-buffer) ----
load_tile(0, 0);
for (int ti = 0; ti <= t_end; ti++) {
int buf = ti & 1;
cp_async_wait_group<0>();
__syncthreads();
if (ti < t_end) load_tile(ti + 1, (ti + 1) & 1);
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
int kv0 = ti * Traits::BC;
if (!IsCausal || kv0 <= max_kv) {
float Sacc[Traits::NC8][4];
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
#pragma unroll
for (int n8 = 0; n8 < Traits::NC8; n8++)
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
int maxc0 = IsCausal ? min(seq_len, causal_off + qr0 + 1)
: seq_len;
int maxc1 = IsCausal ? min(seq_len, causal_off + qr1 + 1)
: seq_len;
// HasMask: mask[batch, q_head, qi, kc] — kc is request-local.
mma_softmax_tile<Traits, HasMask>(kv0, maxc0, maxc1,
qr0, qr1,
p.mask_b_stride, p.mask_h_stride,
p.mask_q_stride,
req_b, q_head,
p.mask,
Sacc, Oacc, m0, m1, l0, l1, lane);
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
}
}
// ---- write output: packed bf16x2 stores ----
float rl0 = (l0 > 1e-20f) ? (1.0f / l0) : 0.0f;
float rl1 = (l1 > 1e-20f) ? (1.0f / l1) : 0.0f;
const int o_base = p.qo_indptr[req_b] * p.q_stride_l + q_head * p.q_stride_h;
#pragma unroll
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
int d = dn8 * 8 + 2 * tid4;
if (qr0 < q_len) {
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][0] * rl0,
Oacc[dn8][1] * rl0);
*reinterpret_cast<__nv_bfloat162*>(
&p.o[o_base + qr0 * p.q_stride_l + d * p.q_stride_d]) = v;
}
if (qr1 < q_len) {
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1,
Oacc[dn8][3] * rl1);
*reinterpret_cast<__nv_bfloat162*>(
&p.o[o_base + qr1 * p.q_stride_l + d * p.q_stride_d]) = v;
}
}
}
+7 -3
View File
@@ -10,15 +10,19 @@ torch::Tensor attn_prefill(
double scale,
int64_t layout
) {
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
auto stream = at::cuda::getCurrentCUDAStream();
AttentionParams<bf16> 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");
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();
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;
}
@@ -30,6 +34,6 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
py::arg("mask") = py::none(),
py::arg("causal_offset") = -1,
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)");
}
+1 -1
View File
@@ -64,7 +64,7 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
// KV: stride-based base
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
int mask_batch_base = batch * p.mask_b_stride;
int mask_batch_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
int tiles = (p.kv_len + P_BC - 1) / P_BC;
int tt = G * ROWS;
int lid = row * G + gpos;
+2 -2
View File
@@ -114,8 +114,8 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
: p.kv_len;
mma_softmax_tile<Traits, HasMask>(kv0, maxc0, maxc1,
qr0, qr1,
p.mask_b_stride, p.mask_q_stride,
batch,
p.mask_b_stride, p.mask_h_stride, p.mask_q_stride,
batch, q_head,
p.mask,
Sacc, Oacc, m0, m1, l0, l1, lane);
+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)"
);
}
-185
View File
@@ -1,185 +0,0 @@
/*
Pure-C test — uses shared dispatcher.
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_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);
}
// Warmed-up, CUDA-event timed sweep over the production decode MMA path.
static void bench() {
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 = 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;
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); }); };
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);
free_scratch(sc);
}
}
static int run_test(int B, int Hq, int Hk, int sl, int D, int causal) {
int gs = Hq / Hk;
printf("=== B=%d Hq=%d Hk=%d seq=%d D=%d gs=%d causal=%d ===\n",
B,Hq,Hk,sl,D,gs,causal);
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); });
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, 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-8f);
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; }
}
printf("kernel: %.3f ms max_abs_err: %.6e max_rel_err: %.6e %s\n\n",
kms, max_abs_err, max_rel_err, pass?"PASS":"FAIL");
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;
}
int main() {
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]);
int fail = 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], causal = configs[ci][5];
fail += run_test(B, Hq, Hk, sl, D, causal);
if (fail) break;
}
if (fail) {
printf("FAILED\n");
return fail;
}
printf("All tests passed!\n");
bench();
return 0;
}
-308
View File
@@ -1,308 +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_dispatchers.cuh"
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 int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int causal, int seed) {
printf("B=%d Hq=%d Hkv=%d kv_len=%d page_sz=%d head_dim=%d causal=%d ... ",
B, Hq, Hkv, kv_len, page_size, HEAD_DIM, causal);
fflush(stdout);
int max_pages = (kv_len + page_size - 1) / page_size;
int n_phys_pages = B * max_pages;
int max_splits = 32;
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);
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;
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_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, causal ? 0 : -1);
PagedAttentionParams<bf16> 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 = causal ? 0 : -1;
set_default_paged_strides(p);
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
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;
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(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_abs_err = 0.0f, max_rel_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_abs_err) { max_abs_err = e; bad_idx = i; }
float rel = e / fmaxf(fabsf(h_o_ref[i]), 1e-8f);
if (rel > max_rel_err) max_rel_err = rel;
}
const float atol = 0.01f, rtol = 0.01f;
bool pass = true;
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
float e = fabsf(h_o_paged[i] - h_o_ref[i]);
if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; }
}
if (pass) {
printf("PASS (max_abs_err=%.4e max_rel_err=%.4e)\n", max_abs_err, max_rel_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 max_rel_err=%.4e at [%d,%d,%d]: ref=%.4f got=%.4f)\n",
max_abs_err, max_rel_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_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, causal, seed;
};
static const TestCase TESTS[] = {
{128, 1, 1, 1, 8, 128, 0, 1},
{128, 1, 4, 4, 128, 128, 0, 2},
{128, 2, 4, 4, 256, 128, 0, 3},
{128, 1, 4, 1, 64, 64, 0, 4},
{128, 1, 8, 2, 64, 128, 0, 5},
{128, 2, 16, 4, 128, 128, 0, 6},
{64, 1, 4, 2, 32, 128, 0, 7},
{256, 1, 2, 1, 16, 128, 0, 8},
{32, 1, 4, 2, 32, 64, 0, 9},
{128, 3, 8, 2, 256, 128, 0, 10},
{128, 2, 32, 8, 512, 128, 0, 11},
{128, 1, 16, 2, 256, 128, 0, 12},
{128, 2, 32, 4, 512, 128, 0, 13},
{128, 2, 8, 2, 128, 128, 1, 14}, // causal
};
static int dispatch_test(const TestCase& tc) {
int r = 0;
dispatch_by_head_dim(tc.head_dim, [&]<int D>() {
r = run_test<D>(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.causal, tc.seed);
});
return r;
}
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;
int max_splits = 32;
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);
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);
PagedAttentionParams<bf16> 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.scale = 1.0f / sqrtf((float)HEAD_DIM);
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 = [&]() {
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(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
PagedAttentionParams<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);
PagedAttentionParams<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
PagedAttentionParams<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);
PagedAttentionParams<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);
PagedAttentionParams<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);
PagedAttentionParams<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;
}
-169
View File
@@ -1,169 +0,0 @@
/*
Pure-C test — uses shared dispatcher.
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_dispatchers.cuh"
// Warmed-up, CUDA-event timed throughput sweep over the production MMA path.
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;
auto launch = [&]() { dispatch_by_head_dim(D, [&]<int H>() { dispatch_prefill<H>(p); }); };
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;
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);
}
}
static int run_test(int B, int Hq, int Hk, int ql, int kl, int D, int causal) {
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_by_head_dim(D, [&]<int H>() { dispatch_prefill<H>(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_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-8f);
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; }
}
printf("kernel: %.3f ms max_abs_err: %.6e max_rel_err: %.6e %s\n\n",
kms, max_abs_err, max_rel_err, pass?"PASS":"FAIL");
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp;
return pass ? 0 : 1;
}
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]);
int fail = 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];
fail += run_test(B, Hq, Hk, ql, kl, D, causal);
if (fail) break;
}
if (fail) {
printf("FAILED\n");
return fail;
}
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-8f);
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-8f);
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;
}
+25 -9
View File
@@ -29,19 +29,18 @@ inline double now_ms() {
struct BenchResult {
float ms;
double gbps;
double tflops;
};
template <typename Fn>
BenchResult bench_kernel(Fn launch, int warmup, int iters,
double flops, double bytes) {
double flops) {
for (int i = 0; i < warmup; i++) launch();
cudaDeviceSynchronize();
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
printf("CUDA error before bench: %s\n", cudaGetErrorString(err));
return {0, 0, 0};
return {0, 0};
}
cudaEvent_t s, e;
@@ -52,19 +51,33 @@ BenchResult bench_kernel(Fn launch, int warmup, int iters,
float ms = 0; cudaEventElapsedTime(&ms, s, e); ms /= iters;
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() {
printf("%-46s | %10s | %10s | %10s\n",
"config", "latency", "bandwidth", "throughput");
printf("%-46s | %10s | %10s\n",
"config", "latency", "TFLOP/s");
printf("---------------------------------------------------------------"
"----------------------------\n");
}
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",
cfg, r.ms, r.gbps, r.tflops);
printf("%-46s | %7.4f ms | %6.2f\n",
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>
@@ -103,6 +116,7 @@ inline void set_default_strides(P& p) {
p.kv_stride_l = p.head_dim;
p.kv_stride_d = 1;
p.mask_b_stride = p.kv_len;
p.mask_h_stride = 0;
p.mask_q_stride = 0;
}
@@ -114,6 +128,7 @@ inline void set_default_paged_strides(P& p) {
p.q_stride_l = p.head_dim;
p.q_stride_d = 1;
p.mask_b_stride = p.kv_len;
p.mask_h_stride = 0;
p.mask_q_stride = 0;
}
@@ -133,9 +148,10 @@ static void cpu_attention_ref(
float scale = 1.0f / sqrtf((float)D);
int n_rep = Hq / Hk;
for (int b = 0; b < B; b++) {
#pragma omp parallel for collapse(2) schedule(dynamic)
for (int h = 0; h < Hq; h++) {
int kv_h = h / n_rep;
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 = kv_len;
+4
View File
@@ -3,6 +3,8 @@ services:
build:
context: .
dockerfile: Dockerfile
args:
CUDA_TAG: ${CUDA_TAG:-cu128}
user: "${UID:-1000}:${GID:-1000}"
ports:
- "8000:8000"
@@ -29,6 +31,8 @@ services:
build:
context: .
dockerfile: Dockerfile
args:
CUDA_TAG: ${CUDA_TAG:-cu128}
user: "${UID:-1000}:${GID:-1000}"
ports:
- "8000:8000"
@@ -1,9 +1,9 @@
<div align="center">
<img src="../images/logo.png" width="auto" alt="Logo">
<img src="./images/logo.png" width="auto" alt="Logo">
<div>
<a href="../../README.md">English</a> •
<a href="../README.md">English</a> •
<a href="#chinese">中文</a>
</div>
@@ -23,7 +23,7 @@
<br>
<div align="center">
<a href="../../README.md">English</a> •
<a href="../README.md">English</a> •
<a href="#chinese">中文</a> •
<a href="https://github.com/ViperEkura/AstrAI/issues">问题追踪</a> •
<a href="https://github.com/ViperEkura/AstrAI/discussions">讨论区</a> •
@@ -62,6 +62,8 @@
**1. 安装**
AstrAI 需要 Python 3.12+,并精确固定 PyTorch 版本为 `2.11.0`。训练、`scripts/tools/generate.py`、生成式评估和生成演示需要 CUDA;CPU 支持仅适用于提供明确 CPU 设备路径的组件,例如 HTTP 服务和直接打分评估。
```bash
git clone https://github.com/ViperEkura/AstrAI.git
cd AstrAI
@@ -138,7 +140,7 @@ curl http://localhost:8000/v1/chat/completions \
# 下载模型权重(运行演示前必需)
python scripts/demo/download.py # model → params/
# 交互式流式聊天(多轮对话,保持历史记录
# 单轮交互式流式提示循环(不保留对话历史
python scripts/demo/stream_chat.py
# 在 >> 后输入消息,输入 !exit 退出
@@ -189,7 +191,7 @@ docker run --gpus all -v /path/to/data:/data -it astrai:latest
# Docker ComposeGPU,默认)
docker compose up -d
# Docker Compose(仅 CPU
# Docker Compose CPU 服务配置(不支持仅限 CUDA 的生成脚本和演示
docker compose --profile cpu up -d
```
@@ -219,22 +221,27 @@ curl -X POST http://localhost:8000/v1/messages \
curl http://localhost:8000/health
```
SSE 流式格式、错误码和统计端点详见[推理文档](./inference.md)。
SSE 流式格式、错误码和统计端点详见[推理文档](guides/inference.md)。
### 文档
| 文档 | 说明 |
|------|------|
| [CLI 参考](./params.md) | 所有 CLI 工具参数(训练、服务、生成、预处理) |
| [架构文档](./architecture.md) | 系统架构、类图与设计模式 |
| [训练文档](./training.md) | 训练循环、策略与公式 |
| [推理文档](./inference.md) | KVCache、连续批处理、采样与 HTTP API |
| [数据流程](./dataflow.md) | 数据管道、存储后端与数据集架构 |
| [数据预处理](./preprocessing.md) | 声明式 JSON 驱动数据预处理 |
| [快速上手](./get-started.md) | 安装与快速入门 |
| [CLI 参考](./guides/params.md) | 所有 CLI 工具参数(训练、服务、生成、预处理) |
| [数据预处理](./guides/preprocessing.md) | 声明式 JSON 驱动数据预处理 |
| [训练文档](./guides/training.md) | 训练循环、策略与公式 |
| [推理文档](./guides/inference.md) | KVCache、连续批处理、采样与 HTTP API |
| [评估文档](./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 本仓库。
2. 创建功能分支。
@@ -251,10 +258,10 @@ SSE 流式格式、错误码和统计端点详见[推理文档](./inference.md)
### 许可证
本项目采用 [GPL-3.0 许可证](../../LICENSE)。
本项目采用 [GPL-3.0 许可证](../LICENSE)。
---
<div align="center">
<em>专为高性能与易用性设计的轻量级 Transformer 框架。</em>
</div>
</div>
@@ -4,7 +4,7 @@
- [Class Diagram](#class-diagram) — Full Mermaid class diagram across 10+ namespaces
- [Module Overview](#module-overview) — Component inventory per module
- [Design Patterns](#design-patterns) — 13 documented patterns with classes
- [Design Patterns](#design-patterns) — 15 documented patterns with classes
- [Core Relationships](#core-relationships) — 11 key inter-component relationships
## Class Diagram
@@ -49,6 +49,11 @@ classDiagram
+Optional[int] n_shared_experts
+Optional[int] n_activated_experts
+Optional[str] topk_method
+Optional[int] moe_intermediate_size
+Optional[int] shared_expert_intermediate_size
+bool norm_topk_prob
+int decoder_sparse_step
+Optional[List[int]] mlp_only_layers
}
class EncoderConfig {
@@ -63,6 +68,7 @@ classDiagram
+Optional[int] num_attention_heads
+Optional[int] num_key_value_heads
+Optional[bool] use_qk_norm
+Optional[bool] use_gated_attention
+str ffn_type
+Optional[dict] rope_scaling
+Optional[str] pooling_type
@@ -114,22 +120,25 @@ classDiagram
+Dataset dataset
+Callable optimizer_fn
+Callable scheduler_fn
+Optional[str] optimizer_name
+Dict[str, Any] optimizer_hyperparameters
+int n_epoch
+int batch_per_device
+int grad_accum_steps
+Optional[float] max_grad_norm
+list gradient_checkpointing_modules
+Optional[str] compile_mode
+int start_epoch
+int start_samples
+str ckpt_dir
+int ckpt_interval
+str log_dir
+List[str] metrics
+Optional[LoRAConfig] lora
+int random_seed
+int num_workers
+Optional[int] prefetch_factor
+bool pin_memory
+Optional[Callable] collate_fn
+int nprocs
+str backend
+str master_addr
@@ -140,6 +149,7 @@ classDiagram
+Optional[float] val_split
+int val_step
+float neftune_alpha
+float moe_aux_loss_coef
+str parallel_mode
+int rollout_interval
+float rollout_temperature
@@ -149,7 +159,6 @@ classDiagram
+Optional[Callable] reward_model_fn
+dict executor_kwargs
+dict extra_kwargs
+validate()
}
}
@@ -205,10 +214,6 @@ classDiagram
-_fetch_record_key(key, index) Tensor
}
class H5Store {
+load(path)
}
class MmapStore {
+List _mmap_refs
+load(path)
@@ -260,11 +265,15 @@ classDiagram
}
namespace model {
class AutoModel {
+BaseModelConfig config
class ModelFactory {
+Dict _entries
+register(name) decorator
+get_component_class(name) Type
}
class AutoModel {
<<nn.Module>>
+BaseModelConfig config
+from_pretrained(path, disable_random_init, strict) nn.Module
+save_pretrained(save_directory)
+to(*args, **kwargs) Self
@@ -277,7 +286,7 @@ classDiagram
+ModuleList layers
+RMSNorm norm
+Linear lm_head
+forward(input_ids, input_mask, paged_cache, position_ids) Dict[str, Tensor]
+forward(input_ids, input_mask, kv_cache, position_ids) Dict[str, Tensor]
+load_state_dict(state_dict, strict, assign)
+state_dict()
}
@@ -299,7 +308,13 @@ classDiagram
+RMSNorm input_norm
+nn.Module mlp # MLP or DeepSeekMoE via FFNFactory
+RMSNorm post_attention_norm
+forward(x, rotary_emb, attention_mask, paged_cache) Tensor
+forward(x, rotary_emb, attention_mask, kv_cache, is_causal) DecoderOutput
}
class DecoderOutput {
<<TypedDict>>
+Tensor hidden_states
+Optional[Tensor] aux_loss
}
class GQA {
@@ -314,7 +329,7 @@ classDiagram
+Linear q_proj, k_proj, v_proj, o_proj
+Linear gate # only if use_gated_attention
+RMSNorm q_norm, k_norm # only if use_qk_norm
+forward(x, rotary_emb, attn_mask, paged_cache) Tensor
+forward(x, rotary_emb, attn_mask, kv_cache, is_causal) Tensor
}
class MLA {
@@ -334,12 +349,18 @@ classDiagram
+Linear gate # only if use_gated_attention
+RMSNorm kv_norm
+RMSNorm q_norm, k_norm # only if use_qk_norm
+forward(x, rotary_emb, attn_mask, paged_cache) Tensor
+forward(x, rotary_emb, attn_mask, kv_cache, is_causal) Tensor
}
class MLP {
+Linear up, gate, down
+forward(x) Tensor
+forward(x) FFNOutput
}
class FFNOutput {
<<TypedDict>>
+Tensor hidden_states
+Optional[Tensor] aux_loss
}
class DeepSeekMoE {
@@ -351,7 +372,7 @@ classDiagram
+Linear router
+ModuleList shared_experts
+ModuleList routed_experts
+forward(x) Tensor
+forward(x) FFNOutput
}
class AttnFactory {
@@ -380,6 +401,7 @@ classDiagram
+int max_len
+float base
+Optional[Dict] rope_scaling
+Tensor freqs_cis
+forward(x, position_ids=None) Tensor
}
@@ -484,10 +506,6 @@ classDiagram
+save(output_dir, domain, shard_idx, tensors)
}
class H5Writer {
+save(output_dir, domain, shard_idx, tensors)
}
class Pipeline {
+PipelineConfig config
+List[str] paths
@@ -557,7 +575,7 @@ classDiagram
class Trainer {
+TrainConfig train_config
+List[TrainCallback] callbacks
+train(resume_dir)
+train(param_path=None, resume=False)
-_get_default_callbacks() List[TrainCallback]
}
@@ -574,13 +592,17 @@ classDiagram
+int epoch
+int consumed_samples
+float loss
+float grad_norm
+Dict[str, float] metrics
+Optional[float] grad_norm
+GradSNRTracker grad_snr_tracker
+DataLoader val_dataloader
+float val_loss
+Optional[float] val_loss
+int world_size
+int rank
+dict kwargs
+optimizer_step() int
+stop_requested (property) bool
+optimizer_step (property) int
+request_stop()
}
class TrainContextBuilder {
@@ -592,11 +614,22 @@ classDiagram
class BaseStrategy {
+Callable model
+Optional[BaseExecutor] executor
+Optional[Callable] model_fn
+float moe_aux_loss_coef
+dict extra_kwargs
+str device
+__call__(batch) Tensor
+__call__(batch) LossOutput
+compute_loss(batch) Tensor
+compute_loss_output(batch) LossOutput
+supports_online() bool
+set_rollout_runner(runner)
+prepare_from_rollout(result) Dict
+on_optimizer_step()
}
class LossOutput {
<<TypedDict>>
+Tensor loss
+Dict[str, float] metrics
}
class StrategyFactory {
@@ -634,9 +667,12 @@ classDiagram
class RawRollout {
+Tensor prompts
+Tensor prompt_mask
+Tensor responses
+Tensor response_mask
+Tensor logprobs_old
+List[str] prompt_texts
+List[List[str]] response_texts
}
class RolloutResult {
@@ -645,10 +681,18 @@ classDiagram
class BaseRewardModel {
<<abstract>>
+score(prompts, responses) Tensor
+score(List[str] prompts, List[List[str]] responses) Tensor
}
class RolloutGenerator {
+InferenceScheduler scheduler
+int max_tokens
+int group_size
+float temperature
+int top_k
+float top_p
+float frequency_penalty
+int rep_window
+generate(batch) RawRollout
}
@@ -738,7 +782,7 @@ classDiagram
}
class MetricCallback {
+Path log_dir
+Path ckpt_dir
+int save_interval
+List[str] metrics
+int val_step
@@ -762,9 +806,9 @@ classDiagram
+nn.Module model
+AutoTokenizer tokenizer
+InferenceScheduler scheduler
+generate(prompt, stream, max_tokens, temperature, top_p, top_k) Union[Generator, str, List[str]]
+generate(prompt, stream, max_tokens, temperature, top_p, top_k, frequency_penalty, rep_window) Union[Generator, str, List[str]]
+generate_with_request(request) Union[Generator, str, List[str]]
+generate_async(prompt, max_tokens, temperature, top_p, top_k) AsyncGenerator
+generate_async(prompt, max_tokens, temperature, top_p, top_k, frequency_penalty, rep_window) AsyncGenerator
+get_stats() Dict
+shutdown()
}
@@ -772,18 +816,18 @@ classDiagram
class Executor {
+AutoModel model
+AutoTokenizer tokenizer
+KVCache page_cache
+PagePool kv_cache
+Optional[str] device
+Optional[torch.dtype] dtype
+execute_prefill(tasks, prompt_len, start_pos)
+execute_decode(tasks) List[int]
+execute_decode(tasks, return_logprobs=False) Union[List[int], List[Tuple[int, float]]]
}
class InferenceScheduler {
+KVCache _page_cache
+PagePool _cache
+Executor _executor
+TaskManager _task_mgr
+bool _running
+Event _stop_event
+Thread _loop_thread
+int max_seq_len
+str device
@@ -793,6 +837,7 @@ classDiagram
+start()
+stop()
+get_stats() Dict
+run_batch(prompt_ids_list, max_tokens, temperature, top_p, top_k, frequency_penalty, rep_window, return_logprobs) Union[List[List[int]], List[Tuple[List[int], List[float]]]]
}
class Allocator {
@@ -814,85 +859,48 @@ classDiagram
+record(page_idx, token_ids, logical_page_idx)
}
class PagePool {
-Allocator _alloc
-PrefixCache _prefix
+alloc() int
+free(idx)
+inc_ref(idx)
+lookup(token_ids) List[int]
+record(page_idx, token_ids, logical_page_idx)
class KVStorage {
+int size
+Tensor k_buffer
+Tensor v_buffer
+get_key_buffer(layer_id) Tensor
+get_value_buffer(layer_id) Tensor
+set_kv_buffer(layer_id, loc, k, v)
}
class Storage {
+int page_size
+Tensor k_cache
+Tensor v_cache
+write(layer_id, page_table, start_pos, k, v)
+gather(layer_id, page_table, total_len) Tuple[Tensor, Tensor]
class ReqToTokenPool {
+int size
+int max_context_len
+Tensor req_to_token
+alloc(num_reqs) List[int]
+free(req_indices)
+write(indices, values)
}
class KVCache {
<<abstract>>
+task_alloc(task_id, prompt_ids) bool
+task_free(task_id)
+task_extend(task_id, pos) bool
+task_cached(task_id) int
+task_record_hashes(task_id, prompt_ids, start_logical_page)
+bind_tasks(task_ids, total_len, device) CacheView
+Tensor k_buffer
+Tensor v_buffer
+Tensor req_to_token
+Tensor req_pool_indices
+Tensor seq_lens
+Tensor out_cache_loc
+int max_len
+Optional[Tensor] kv_indptr
}
class PageCache {
class PagePool {
+int page_size
-PagePool _pool
-Storage _storage
-TaskTable _table
+bool contiguous
-KVStorage _storage
-ReqToTokenPool _req_pool
-Allocator _alloc
-PrefixCache _prefix
+task_alloc(task_id, prompt_ids) bool
+task_free(task_id)
+task_extend(task_id, pos) bool
+task_cached(task_id) int
+task_record_hashes(task_id, prompt_ids, start_logical_page)
+bind_tasks(task_ids, total_len, device) PageCacheView
}
class ContiguousCache {
+int max_seq_len
+Tensor k, v
+task_alloc(task_id, prompt_ids) bool
+task_free(task_id)
+task_extend(task_id, pos) bool
+bind_tasks(task_ids, total_len, device) ContiguousCacheView
}
class CacheView {
<<abstract>>
+write(layer_id, k, v)
+gather(layer_id) Tuple[Tensor, Tensor]
}
class PageCacheView {
-Storage _storage
+Tensor _page_table
+int _total_len
+write(layer_id, k, v)
+gather(layer_id) Tuple[Tensor, Tensor]
}
class ContiguousCacheView {
-ContiguousCache _cache
+Tensor _batch_indices
+int _total_len
+write(layer_id, k, v)
+gather(layer_id) Tuple[Tensor, Tensor]
}
class TaskTable {
+set(task_id, page_table, cached)
+get(task_id) List[int]
+get_cached(task_id) int
+get_ref(task_id) List[int]
+pop(task_id) Tuple[List[int], int]
+table_tensor(task_ids, device) Tensor
+bind_tasks(task_ids, seq_lens, device, start_pos) KVCache
}
class Task {
@@ -902,6 +910,8 @@ classDiagram
+float temperature
+float top_p
+int top_k
+float frequency_penalty
+int rep_window
+TaskStatus status
+List output_ids
+int input_tokens
@@ -924,7 +934,6 @@ classDiagram
+AutoTokenizer tokenizer
+int max_batch_size
+int max_seq_len
+int max_prompt_len
+Deque waiting_queue
+List active_tasks
+add_task(prompt, max_tokens, temperature, top_p, top_k, stream_callback) str
@@ -948,27 +957,29 @@ classDiagram
+float top_p
+float temperature
+Optional[int] max_tokens
+float frequency_penalty
+int rep_window
+bool stream
}
class BaseSamplingStrategy {
<<abstract>>
+apply(logits, filter_value) Tensor
+apply(logits, filter_value, input_ids, input_mask) Tensor
}
class TemperatureStrategy {
+float temperature
+apply(logits, filter_value) Tensor
+apply(logits, filter_value, input_ids, input_mask) Tensor
}
class TopKStrategy {
+int top_k
+apply(logits, filter_value) Tensor
+apply(logits, filter_value, input_ids, input_mask) Tensor
}
class TopPStrategy {
+float top_p
+apply(logits, filter_value) Tensor
+apply(logits, filter_value, input_ids, input_mask) Tensor
}
class FrequencyPenaltyStrategy {
@@ -978,8 +989,8 @@ classDiagram
class SamplingPipeline {
+List[BaseSamplingStrategy] strategies
+apply(logits, filter_value) Tensor
+sample(logits, filter_value) Tensor
+apply(logits, filter_value, input_ids, input_mask) Tensor
+sample(logits, filter_value, input_ids, input_mask, return_logprobs) Union[Tensor, Tuple[Tensor, Tensor]]
}
class StreamDecoder {
@@ -1054,7 +1065,7 @@ classDiagram
<<abstract>>
+prepare(request, engine) Tuple[str, GenContext, List[str]]
+format_stream_start(ctx) List[str]
+format_chunk(token) List[str]
+format_chunk(token, **kwargs) List[str]
+format_stream_end(ctx, stop) List[str]
+format_response(ctx, content, stop) Dict
}
@@ -1062,7 +1073,7 @@ classDiagram
class OpenAIResponseBuilder {
+prepare(request, engine) Tuple
+format_stream_start(ctx) List[str]
+format_chunk(token) List[str]
+format_chunk(token, **kwargs) List[str]
+format_stream_end(ctx, stop) List[str]
+format_response(ctx, content, stop) Dict
}
@@ -1070,7 +1081,7 @@ classDiagram
class AnthropicResponseBuilder {
+prepare(request, engine) Tuple
+format_stream_start(ctx) List[str]
+format_chunk(token) List[str]
+format_chunk(token, **kwargs) List[str]
+format_stream_end(ctx, stop) List[str]
+format_response(ctx, content, stop) Dict
}
@@ -1178,10 +1189,13 @@ classDiagram
class BaseExecutor {
+GradientState gradient_state
+prepare(model_fn, optimizer_fn, scheduler_fn, before_wrap) tuple
+prepare(model_fn, optimizer_fn, scheduler_fn, before_wrap, after_wrap) tuple
+accumulate(model) context manager
+backward(loss)
+unwrap_model(model) dict
+checkpoint_context(model) context manager
+clip_grad_norm(model, max_norm) float
+use_distributed (property) bool
+sync_gradients (property) bool
+grad_accum_steps (property) int
}
@@ -1196,14 +1210,10 @@ classDiagram
}
class FSDPExecutor {
-_prepare_model(model) nn.Module
+unwrap_model(model) dict
}
class FSDP2Executor {
-_prepare_model(model) nn.Module
-_no_sync(model) context manager
+unwrap_model(model) dict
+unwrap_model(model) Optional[dict]
+clip_grad_norm(model, max_norm) float
}
class ExecutorFactory {
@@ -1212,33 +1222,6 @@ classDiagram
+create(parallel_mode, **kwargs) BaseExecutor
}
class ParallelModel {
+dist.ProcessGroup process_group
+int rank
+int world_size
}
class ColumnParallelLinear {
+int in_features
+int out_features
+int out_features_per_rank
+bool gather_results
+Parameter weight
+Optional[Parameter] bias
+forward(x) Tensor
+load_state_dict(state_dict)
}
class RowParallelLinear {
+int in_features
+int out_features
+int in_features_per_rank
+bool reduce_results
+Parameter weight
+Optional[Parameter] bias
+forward(x) Tensor
+load_state_dict(state_dict)
}
}
%% Relationships — UML notation: <|-- generalization, *-- composition, o-- aggregation, --> association, ..> dependency
@@ -1260,11 +1243,8 @@ classDiagram
BaseDataset <|-- SFTDataset
BaseDataset <|-- DPODataset
BaseDataset <|-- GRPODataset
Store <|-- H5Store
Store <|-- MmapStore
Store <|-- JsonlStore
H5Store --|> Streamable
H5Store --|> Recordable
MmapStore --|> Streamable
MmapStore --|> Recordable
JsonlStore --|> Streamable
@@ -1273,8 +1253,6 @@ classDiagram
BaseSamplingStrategy <|-- TopKStrategy
BaseSamplingStrategy <|-- TopPStrategy
BaseSamplingStrategy <|-- FrequencyPenaltyStrategy
ParallelModel <|-- RowParallelLinear
ParallelModel <|-- ColumnParallelLinear
AutoModel <|-- AutoRegressiveLM
AutoModel <|-- EmbeddingEncoder
BaseConfig <|-- BaseModelConfig
@@ -1285,7 +1263,7 @@ classDiagram
BaseConfig <|-- PipelineConfig
BaseModelConfig <|-- AutoRegressiveLMConfig
BaseModelConfig <|-- EncoderConfig
BaseFactory <|-- AutoModel
BaseFactory <|-- ModelFactory
BaseFactory <|-- AttnFactory
BaseFactory <|-- FFNFactory
BaseFactory <|-- DatasetFactory
@@ -1303,7 +1281,6 @@ classDiagram
BaseExecutor <|-- NoneExecutor
BaseExecutor <|-- DDPExecutor
BaseExecutor <|-- FSDPExecutor
BaseExecutor <|-- FSDP2Executor
ResponseBuilder <|-- OpenAIResponseBuilder
ResponseBuilder <|-- AnthropicResponseBuilder
BaseToolParser <|-- SimpleJsonToolParser
@@ -1317,21 +1294,16 @@ classDiagram
PositionIdStrategy <|-- DocResetPositionId
PositionIdStrategy <|-- ContinuousPositionId
StoreWriter <|-- BinWriter
StoreWriter <|-- H5Writer
RawRollout <|-- RolloutResult
LaunchStrategy <|-- TorchrunStrategy
LaunchStrategy <|-- LocalStrategy
KVCache <|-- PageCache
KVCache <|-- ContiguousCache
CacheView <|-- PageCacheView
CacheView <|-- ContiguousCacheView
%% --- Composition (strong ownership, part destroyed with whole) ---
PageCache *-- PagePool
PageCache *-- Storage
PageCache *-- TaskTable
PagePool *-- KVStorage
PagePool *-- ReqToTokenPool
PagePool *-- Allocator
PagePool *-- PrefixCache
InferenceEngine *-- InferenceScheduler
InferenceScheduler *-- KVCache
InferenceScheduler *-- PagePool
InferenceScheduler *-- Executor
InferenceScheduler *-- TaskManager
AutoRegressiveLM *-- DecoderBlock
@@ -1352,15 +1324,11 @@ classDiagram
%% --- Aggregation (weak ownership) ---
AutoModel o-- BaseModelConfig
AutoTokenizer o-- ChatTemplate
PagePool o-- Allocator
PagePool o-- PrefixCache
Trainer o-- TrainCallback
TrainContext o-- BaseStrategy
TrainContext o-- BaseScheduler
TrainContext o-- Checkpoint
TrainContext o-- BaseExecutor
PageCacheView o-- Storage
ContiguousCacheView o-- ContiguousCache
SamplingPipeline o-- BaseSamplingStrategy
BaseDataset o-- Store
Pipeline o-- PipelineConfig
@@ -1389,15 +1357,15 @@ classDiagram
FFNFactory ..> DeepSeekMoE : creates
DecoderBlock ..> AttnFactory : uses
DecoderBlock ..> FFNFactory : uses
StoreFactory ..> H5Store : creates
StoreFactory ..> MmapStore : creates
StoreFactory ..> JsonlStore : creates
ConfigFactory ..> AutoRegressiveLMConfig : creates
ConfigFactory ..> EncoderConfig : creates
ModelFactory ..> AutoRegressiveLM : creates
ModelFactory ..> EmbeddingEncoder : creates
ExecutorFactory ..> NoneExecutor : creates
ExecutorFactory ..> DDPExecutor : creates
ExecutorFactory ..> FSDPExecutor : creates
ExecutorFactory ..> FSDP2Executor : creates
ToolParserFactory ..> BaseToolParser : creates
TrainContextBuilder ..> ExecutorFactory : creates
Trainer ..> TrainContextBuilder : uses
@@ -1406,8 +1374,7 @@ classDiagram
TrainContextBuilder ..> RDSampler : creates
Checkpoint ..> Checkpoint : serializes
CheckpointCallback ..> Checkpoint : creates
PageCache ..> PageCacheView : binds
ContiguousCache ..> ContiguousCacheView : binds
PagePool ..> KVCache : binds
InferenceEngine ..> GenerationRequest : uses
InferenceEngine ..> GenerateResult : creates
OpenAIResponseBuilder ..> ChatCompletionRequest : receives
@@ -1438,14 +1405,15 @@ classDiagram
| Module | Components | Description |
|--------|------------|-------------|
| **astrai.config** | BaseConfig, BaseModelConfig, AutoRegressiveLMConfig, EncoderConfig, ConfigFactory, TrainConfig, PipelineConfig, InputConfig, ProcessingConfig, OutputConfig | Configuration management (to_dict/from_dict, to_file/from_file) |
| **astrai.preprocessing** | SectionRenderer, BaseMaskBuilder, MaskBuilderFactory, SectionedMaskBuilder, SingleOutputMaskBuilder, MultiOutputMaskBuilder, Pipeline, TokenizeTransform, PackingStrategy, PackingStrategyFactory, SimplePacking, BFDPacking, BFDSplitPacking, PositionIdStrategy, PositionIdStrategyFactory, NoPositionId, DocResetPositionId, ContinuousPositionId, StoreWriter, StoreWriterFactory, BinWriter, H5Writer | Declarative JSON-driven data preprocessing |
| **astrai.dataset** | BaseDataset, SEQDataset, SFTDataset, DPODataset, GRPODataset, Store, Streamable, Recordable, H5Store, MmapStore, JsonlSource, JsonlStore, StoreFactory, RDSampler, DatasetFactory | Dataset loading and management |
| **astrai.preprocessing** | SectionRenderer, BaseMaskBuilder, MaskBuilderFactory, SectionedMaskBuilder, SingleOutputMaskBuilder, MultiOutputMaskBuilder, Pipeline, TokenizeTransform, PackingStrategy, PackingStrategyFactory, SimplePacking, BFDPacking, BFDSplitPacking, PositionIdStrategy, PositionIdStrategyFactory, NoPositionId, DocResetPositionId, ContinuousPositionId, StoreWriter, StoreWriterFactory, BinWriter | Declarative JSON-driven data preprocessing |
| **astrai.dataset** | BaseDataset, SEQDataset, SFTDataset, DPODataset, GRPODataset, Store, Streamable, Recordable, MmapStore, JsonlSource, JsonlStore, StoreFactory, RDSampler, DatasetFactory | Dataset loading and management |
| **astrai.serialization** | Checkpoint | Model serialization |
| **astrai.model** | AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, LoRAConfig, LoRALinear, RotaryEmbedding, Embedding | Neural network model |
| **astrai.model** | ModelFactory, AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, LoRAConfig, LoRALinear, RotaryEmbedding, Embedding | Neural network model |
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategyGRPOStrategy, StrategyFactory, BaseSchedulerWSDScheduler, SchedulerFactory, TrainCallback(Protocol)MetricCallback, CallbackFactory, RawRollout, RolloutResult, BaseRewardModel, RolloutGenerator, RolloutRunner | Training workflow |
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, KVCacheContiguousCache/PageCache, CacheViewContiguousCacheView/PageCacheView, AllocatorStorage, Task, TaskManager, TaskStatus, StreamDecoder, GenerationRequest, GenerateResult, BaseSamplingStrategySamplingPipeline, FrequencyPenaltyStrategy, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, StopInfo, ChatMessage, FunctionDef, ToolDef, ChatCompletionRequest, AnthropicMessage, MessagesRequest, BaseToolParser, ToolParserFactory, SimpleJsonToolParser | Inference service |
| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, LaunchStrategy, TorchrunStrategy, LocalStrategy, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, FSDPExecutor, FSDP2Executor, GradientState, AccumOptimizer, AccumScheduler, ParallelModel, RowParallelLinear, ColumnParallelLinear | Distributed parallel & gradient accumulation |
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, PagePool, KVStorage, ReqToTokenPool, KVCache, Allocator, PrefixCache, Task, TaskManager, TaskStatus, StreamDecoder, GenerationRequest, GenerateResult, BaseSamplingStrategySamplingPipeline, FrequencyPenaltyStrategy, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, StopInfo, ChatMessage, FunctionDef, ToolDef, ChatCompletionRequest, AnthropicMessage, MessagesRequest, BaseToolParser, ToolParserFactory, SimpleJsonToolParser | Inference service |
| **astrai.extension** | AttentionBackend, TorchNativeBackend, CudaBackend, attn_backend, ATTN_BACKEND, attn_decode, attn_prefill, attn_paged_decode, rotary_emb, apply_rotary_emb, rotary_backend, is_available | CUDA attention + rotary kernels, backend abstraction, auto-dispatch |
| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, LaunchStrategy, TorchrunStrategy, LocalStrategy, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, FSDPExecutor, GradientState, AccumOptimizer, AccumScheduler | Distributed parallel & gradient accumulation |
| **astrai.factory** | BaseFactory | Component registration |
| **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers |
@@ -1453,7 +1421,7 @@ classDiagram
| Pattern | Classes | Purpose |
|---------|---------|---------|
| **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory`, `MaskBuilderFactory`, `StoreWriterFactory`, `PackingStrategyFactory`, `PositionIdStrategyFactory`, `ToolParserFactory` | Decorator-based component creation |
| **Factory** | `ModelFactory`, `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory`, `MaskBuilderFactory`, `StoreWriterFactory`, `PackingStrategyFactory`, `PositionIdStrategyFactory`, `ToolParserFactory` | Decorator-based component creation |
| **Registry** | `BaseFactory` | Component registration |
| **Strategy** | `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy` | Training strategy switching |
| **Strategy (Sampling)** | `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations |
@@ -1462,23 +1430,25 @@ classDiagram
| **Observer** | `TrainCallback`, callback implementations | Training process monitoring |
| **Context** | `TrainContext` | Unified training state bag |
| **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction |
| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor`, `FSDP2Executor` | Gradient accumulation & model distribution |
| **Storage** | `Store`, `H5Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support |
| **Strategy (Attention)** | `AttentionBackend`, `TorchNativeBackend`, `CudaBackend` | Attention computation backend switching via context manager |
| **Auto-dispatch (Rotary)** | `apply_rotary_emb`, `rotary_backend.py`, `rotary_ops.py` | Rotary embedding CUDA kernel auto-dispatch with torch fallback |
| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor` | Gradient accumulation & model distribution |
| **Storage** | `Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support |
| **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching |
| **AutoModel Registry** | `AutoModel`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading |
| **Model Registry** | `ModelFactory`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading |
## Core Relationships
1. **Config → Training**: `TrainConfig` holds `model_fn`, `dataset`, `optimizer_fn`, `scheduler_fn`, `parallel_mode`, `executor_kwargs`
2. **Training Flow**: `Trainer``TrainContextBuilder``TrainContext`, uses `BaseStrategy` for loss, `BaseExecutor` for gradient accumulation + model distribution
3. **Strategy Selection**: `StrategyFactory` creates strategy by `train_type`
4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)``NoneExecutor` / `DDPExecutor` / `FSDPExecutor` / `FSDP2Executor`
5. **Inference Flow**: `InferenceEngine``InferenceScheduler``AutoRegressiveLM`, backed by `KVCache` + `SamplingPipeline`
4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)``NoneExecutor` / `DDPExecutor` / `FSDPExecutor`
5. **Inference Flow**: `InferenceEngine``InferenceScheduler``AutoRegressiveLM`, backed by `PagePool` + `KVCache` + `SamplingPipeline`. Attention backend selected via `attn_backend()` context manager (`TorchNativeBackend` default, `CudaBackend` for CUDA kernels). Rotary embedding auto-dispatches to CUDA kernel when available (inference mode), else torch complex multiply (training).
6. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP
7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (H5Store/MmapStore/JsonlStore) loads data with explicit `_length` and multi-segment `_data`
8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata (rank-0 only), extra state saved as `{key}.pt`
7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (`MmapStore`/`JsonlStore`) loads data with explicit `_length` and multi-segment `_data`
8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata; `CheckpointCallback` performs rank-0 training saves, with extra state saved as `{key}.pt`
9. **Scheduler**: `SchedulerFactory` creates `CosineScheduler`/`SGDRScheduler`/`WSDScheduler`
10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops
11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers
> Document Update Time: 2026-07-20
> Document Update Time: 2026-08-02
+173
View File
@@ -0,0 +1,173 @@
# 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 |
| `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+) |
| Paged split-KV MMA decode | `attn_paged_decode_split_kv_mma.cuh` | Paged cache + split-KV + MMA |
### 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
```
### Architecture flags
`csrc/build.py` auto-detects the GPU compute capability and generates the appropriate `nvcc` gencode flag:
- **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
```
NVCC_FLAGS = -O3 --expt-relaxed-constexpr --use_fast_math
--ptxas-options=-O3,-v --extra-device-vectorization --threads=8
```
The `REGISTRY` in `csrc/build.py` lists all registered kernels (currently 4). Each entry maps a kernel name to its source files and build flags.
## 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_prefill`
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/
├── build.py # Build system: REGISTRY, _arch_flags, nvcc flags
├── kernels/
│ ├── attn_common.h # Shared attention params (AttentionParams, PagedAttentionParams)
│ ├── attn_decode.cu # Basic decode kernel (registered)
│ ├── attn_prefill.cu # Basic prefill kernel (registered)
│ ├── attn_paged_decode.cu # Paged decode kernel (registered)
│ ├── rotary_emb.cu # Fused rotary embedding kernel (registered)
│ ├── attn_decode_split_kv.cuh # Split-KV variant
│ ├── attn_decode_split_kv_mma.cuh # Split-KV + MMA variant
│ ├── attn_prefill_split_q.cuh # Split-Q variant
│ ├── attn_prefill_split_q_mma.cuh # Split-Q + MMA variant
│ ├── attn_paged_decode_split_kv.cuh # Paged + split-KV variant
│ ├── attn_paged_decode_split_kv_mma.cuh # Paged + split-KV + MMA variant
│ ├── attn_dispatchers.cuh # Kernel dispatch macros
│ ├── 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
+166
View File
@@ -0,0 +1,166 @@
# Data Flow
This document describes the data pipeline: from raw text to model input tensors. For creating preprocessing configs, see [Preprocessing Guide](../guides/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
```
JSON / JSONL Records → Pipeline (mask builder) → Tokenized Tensors
.bin storage
Store.load()
Store.fetch(begin, end, keys)
Dataset.__getitem__(idx)
RDSampler → DataLoader → Training
```
## Data Preparation
The offline `Pipeline` accepts `.jsonl` records and `.json` files containing one
object or a list of objects. It tokenizes them and writes binary shards (`.bin`
plus `meta.json`) with keyed tensor groups. Binary is the only registered output
writer; the pipeline cannot emit JSONL.
### Tokenization
The `Pipeline` reads JSON/JSONL records, applies the mask builder (see
[Preprocessing](../guides/preprocessing.md)), and produces 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
```
For default single-output preprocessing, the stored keys are `sequence` and
`position_ids`, plus `loss_mask` when masking is required. Packing is supported
for single-output data with a `sequence` key. Shard flushing counts the primary
flat sequence for each record: `sequence` in single-output mode, otherwise the
first flat source output.
The exact shard `meta.json` schema is a top-level mapping from key to tensor
metadata. It does not contain a storage-format or total-token field:
```json
{
"sequence": {"shape": [123456], "dtype": "int32"},
"loss_mask": {"shape": [123456], "dtype": "bool"},
"position_ids": {"shape": [123456], "dtype": "int32"}
}
```
Record-aware binary data may also include `"offsets": [0, ...]` inside a key's
metadata, but the preprocessing `BinWriter` currently does not write offsets.
### Format Detection
`detect_format(load_path)` inspects the path:
- If `load_path` is a file: `.jsonl` selects `"jsonl"`; other suffixes raise `ValueError`.
- If `load_path` is a directory: any recursive `*.bin` plus a `meta.json` selects `"bin"`; otherwise any recursive `*.jsonl` selects `"jsonl"`.
- Detection does not require `dataset_config.json`; configuration is selected later when `JsonlStore.load()` chooses a transform.
### Store Backends
Storage format is auto-detected by `detect_format()`; backends are dispatched via registry:
```
StoreFactory.create("bin") → MmapStore
StoreFactory.create("jsonl") → JsonlStore
```
Both stores inherit `Store` and compose the `Streamable` and `Recordable`
access methods.
**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**: Reads a `.jsonl` file or the sorted top-level `*.jsonl` files in
a directory. Eager transform selection uses the first available route:
1. An explicit `transform=` argument.
2. `dataset_config.json` in the JSONL directory. It follows `PipelineConfig` and may add `tokenizer_path`; when omitted, the config directory is used.
3. The built-in `messages` transform when `tokenizer_path=` is supplied. It masks system/user turns, trains assistant turns, and emits document-reset position IDs.
Only DPO gets an automatic lazy route from `DatasetFactory`: raw JSONL plus
`tokenizer_path` installs `dpo_processor` and tokenizes each record in
`fetch_record`. GRPO does not currently have an automatic lazy processor.
Eager-loaded stores normalize tensors into `Store._data[Dict[str, List[Tensor]]]` + `Store._cum[Dict[str, List[int]]]` (cumulative lengths for stream indexing) + `Store._offsets[Dict[str, List[int]]]` (per-record offsets for record indexing). Nested JSONL keys such as GRPO `responses`/`masks` are kept as record values and excluded from stream bookkeeping. Lazy DPO instead retains raw records and processes them in `fetch_record`.
## Data Keys by Training Type
| Type | Storage Keys | Access Mode |
|------|-------------|-------------|
| `seq` | `sequence`, `position_ids` by default (`SEQDataset` consumes only `sequence`) | 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`) |
Offline `.bin` output from DPO/GRPO preprocessing is not currently loadable for
training. DPO shards are written without record offsets, while GRPO response
groups are flattened without preserving record/group boundaries. Supported raw
routes are eager JSONL for SEQ/SFT and automatic lazy JSONL for DPO. GRPO
requires a caller-built, already-loaded record store.
## Dataset Architecture
```
DatasetFactory.load(...)
→ detect_format(load_path)
→ optionally build dpo_processor for raw JSONL
→ StoreFactory.create(storage_type, window_size, stride)
→ Store.load(load_path, transform=... or processor=...)
→ DatasetFactory.create(train_type, store=store)
Stream datasets (SEQ/SFT):
BaseDataset.__getitem__(idx)
→ Store.sample_window(idx) → [begin, end)
→ Store.fetch(begin, end, keys) → Tensor / Dict[str, Tensor]
Record datasets (DPO/GRPO):
DPODataset/GRPODataset.__getitem__(idx)
→ Store.fetch_record(idx, keys) → Tensor / Dict[str, Tensor]
```
Class hierarchy: `BaseDataset` is the direct base of `SEQDataset`, `SFTDataset`,
`DPODataset`, and `GRPODataset`. There is no `RecordDataset` class.
`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`).
For raw JSONL, `tokenizer_path` builds the lazy processor only for DPO. For
SEQ/SFT it is forwarded to `JsonlStore` so the built-in eager `messages`
transform can be selected when no `dataset_config.json` exists. GRPO receives no
automatic processor. A pre-built `store` bypasses path, format, tokenizer,
window, and stride setup entirely.
`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 for binary record layouts; otherwise it indexes per-record JSONL tensors directly.
## Sampler
`RDSampler` 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
+250
View File
@@ -0,0 +1,250 @@
# Internals
Mathematical foundations and internal algorithms for AstrAI's training, inference, and preprocessing pipelines. For practical usage guides, see [Training](../guides/training.md), [Inference](../guides/inference.md), and [Preprocessing](../guides/preprocessing.md).
## Contents
- [Autoregression & Causal Masking](#autoregression--causal-masking)
- [Rotary Position Embedding (RoPE)](#rotary-position-embedding-rope)
- [Training Loss Formulas](#training-loss-formulas)
- [Training Loop Internals](#training-loop-internals)
- [Callback Lifecycle](#callback-lifecycle)
- [KV Cache Mathematics](#kv-cache-mathematics)
- [Mask Algorithm Internals](#mask-algorithm-internals)
- [Gradient Accumulation Mechanics](#gradient-accumulation-mechanics)
## Autoregression & Causal Masking
Given a token sequence, the model predicts the probability of the next token. Each generated token is appended to the input and fed back, repeating until an end-of-sequence token or max length.
```
sequence : [[1, 2, 3, 4, 5, 6]]
input_ids: [[1, 2, 3, 4, 5]]
target_ids: [[2, 3, 4, 5, 6]]
```
A lower-triangular causal mask prevents attending to future positions:
```
[[0, -inf, -inf, -inf, -inf],
[0, 0, -inf, -inf, -inf],
[0, 0, 0, -inf, -inf],
[0, 0, 0, 0, -inf],
[0, 0, 0, 0, 0]]
```
This ensures position $i$ can only attend to positions $\leq i$, which is essential for autoregressive generation.
## Rotary Position Embedding (RoPE)
RoPE embeds position into Q/K vectors via complex rotation:
$$ q_i = R_i W_q x_i, \quad k_j = R_j W_k x_j, \quad q_i^T k_j = x_i^T W_q^T R_{i-j} W_k x_j $$
`RotaryEmbedding` pre-computes a complex `freqs_cis` buffer. `forward()` returns
a tensor indexed by `position_ids`. `apply_rotary_emb` applies the rotation:
during training it uses torch complex multiply (autograd-compatible); during
inference it auto-dispatches to a fused CUDA kernel when available. The key
property is that the dot product $q_i^T k_j$ depends only on the relative
position $i - j$, not the absolute positions.
**Critical for inference**: RoPE is applied **before** KV cache write, not after. If applied after caching, position encoding drift occurs because cached K/V would have stale rotation factors.
## Training Loss Formulas
### SEQ (Pre-training)
Next-token cross-entropy with optional label smoothing:
$$ L_{\text{PT}} = -\frac{1}{T}\sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta) $$
### SFT (Supervised Fine-Tuning)
Masked cross-entropy (`ignore_index=-100`) over response tokens only:
$$ L_{\text{SFT}} = -\frac{1}{L}\sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta) $$
Prompt tokens are masked out via `loss_mask`; only response tokens contribute to the loss.
### DPO (Direct Preference Optimization)
Frozen reference model, preference margin via log-ratio:
$$ L_{\text{DPO}} = -\mathbb{E}\left[\log\sigma\left(\beta\log\frac{\pi_\theta(y_w\mid x)}{\pi_{\text{ref}}(y_w\mid x)} - \beta\log\frac{\pi_\theta(y_l\mid x)}{\pi_{\text{ref}}(y_l\mid x)}\right)\right] $$
Parameters: `beta=0.1`, `reduction="sum"`.
### GRPO (Group Relative Policy Optimization)
Token-level PPO with group-normalized advantages:
$$ \text{Advantage}_i = \frac{r_i - \mu}{\sigma + \epsilon} $$
$$ L_{\text{GRPO}} = -\mathbb{E}_t\left[\min\left(\rho_t A,\; \text{clip}\left(\rho_t, 1-\epsilon, 1+\epsilon\right)A\right)\right] + \lambda \cdot \mathbb{E}_t\left[\frac{\pi_{\text{ref}}}{\pi_\theta} - \log\frac{\pi_{\text{ref}}}{\pi_\theta} - 1\right] $$
Where $\rho_t = \pi_\theta(a_t|s_t) / \pi_{\text{old}}(a_t|s_t)$ is the per-token importance sampling ratio. Advantages are derived from scalar per-response rewards, group-normalized, and broadcast across all response tokens. Only response tokens contribute to the loss.
Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`.
### MoE Load Balancing
MoE layers add a differentiable load-balancing term based on mean router probabilities and top-k expert assignment frequency. The training objective is:
$$ L = L_{\text{task}} + \lambda_{\text{MoE}} L_{\text{aux}} $$
`TrainConfig.moe_aux_loss_coef` controls $\lambda_{\text{MoE}}$ (default `0.01`). The unweighted and weighted auxiliary losses are logged separately.
## Training Loop Internals
Two-level loop: **epoch****batch**. Optimizer step fires every `grad_accum_steps` batches.
```
on_train_begin
model.train()
on_epoch_begin
for batch in dataloader:
with executor.accumulate(model):
on_batch_begin
loss_output = strategy(batch)
context.loss = loss_output["loss"].item()
context.metrics = loss_output["metrics"]
stand_loss = loss_output["loss"] / executor.grad_accum_steps
executor.backward(stand_loss)
context.consumed_samples += (
context.config.batch_per_device * context.world_size
)
on_batch_end
if executor.sync_gradients:
on_optimizer_step
optimizer.step()
strategy.on_optimizer_step()
optimizer.zero_grad()
if scheduler:
scheduler.step()
on_epoch_end
on_train_end
```
The loss is divided by `grad_accum_steps` before `backward()`, so accumulated gradients sum to the correct mean.
Strategy metrics are detached and converted to Python `float` values before the
`LossOutput` is returned; only `LossOutput.loss` remains a differentiable tensor.
## Callback Lifecycle
| Hook | Fires | Default callback |
|------|-------|-----------------|
| `on_train_begin` | Before training starts | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` |
| `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` |
| `on_batch_begin` | Every batch | — |
| `on_optimizer_step` | Every accumulation window | `MetricCallback`, `ProgressBarCallback`, `GradientClippingCallback` |
| `on_batch_end` | Every batch | `CheckpointCallback` |
| `on_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` |
| `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` |
| `on_train_end` | Training exits after `on_train_begin` completes (via `finally`) | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` |
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm, rank-0), `gradient_clipping`. The gradient-clipping callback is always registered and always calls `executor.clip_grad_norm()` with the numeric `max_grad_norm` value.
## KV Cache Mathematics
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 $$
The cache stores $k_j$ and $v_j$ for all previous positions. At each decode step, only $q_n$ (the current query) is computed fresh, and attention is computed against the cached K/V.
**RoPE ordering**: RoPE is applied to Q/K **before** writing to the KV cache. This is essential because:
1. The cached K values already contain the rotation for their original positions.
2. The new Q is rotated for its current position.
3. The dot product $q_n^T k_j$ then correctly depends on $n - j$ (relative position).
If RoPE were applied after caching, the rotation factors would be inconsistent between cached and new tokens.
### Cache Architecture
Three-layer separation (SGLang-inspired):
- **KVStorage**: Flat token-level buffers `[n_layers, size, n_kv_heads, head_dim]`.
- **ReqToTokenPool**: Index table `[req_idx, pos] → physical token slot`, shared across all layers.
- **Allocator + PrefixCache**: Paged-mode slot allocation with ref-counting, LRU eviction, and hash-based prefix sharing.
`PagePool` orchestrates all three. In contiguous mode (default), `req_to_token` is a trivial linear mapping. In paged mode, slots are allocated on demand with prefix caching support. `bind_tasks()` returns a `KVCache` dataclass with `kv_indptr`, a prefix-sum index over sequence lengths computed once per step and shared across layers. Attention layers access buffers directly — no methods, no abstraction.
### Attention Backend
Attention computation is decoupled from the model via `AttentionBackend` ABC (`astrai/extension/attention_backend.py`):
- **`TorchNativeBackend`** (default): writes K/V to cache, gathers via `req_to_token` indirect indexing, calls `F.scaled_dot_product_attention`.
- **`CudaBackend`**: decode path uses `attn_paged_decode` with `page_size=1` (the `req_to_token` table serves as the page table, each token slot is a single-token "page"); prefill path gathers K/V then calls `attn_prefill`. Falls back to `TorchNativeBackend` when kernel unavailable.
Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/rotary_backend.py`, which auto-dispatches to the fused CUDA kernel (`rotary_emb.cu`) during inference or torch complex multiply during training (for autograd compatibility). Both attention backends share the same rotary dispatch.
Backend selection is thread-safe via `contextvars`, mirroring `torch.nn.attention.sdpa_kernel`:
```python
from astrai.extension import attn_backend, ATTN_BACKEND
with attn_backend(ATTN_BACKEND.CUDA):
engine.generate("hello")
```
Layout convention: all q/k/v are `[batch, seq_len, n_heads, head_dim]` (blhd). Scale is always `1/sqrt(head_dim)`.
## Mask Algorithm Internals
### Template mode (`template: true`)
1. Prepend BOS token (masked)
2. For each message in the field's array:
1. Render through `chat_template` for that single message
2. Encode rendered text
3. Apply mask rule for the message's role
### Non-template mode
Encode the field value as text. Mask value is 1 (train) or 0 (mask) per the section's `action`.
### Text config detection
When no section uses `template` and all sections have `action: "train"`, the builder omits `loss_mask` from the output — all tokens are trained.
### Position ID strategies
| Mode | Behavior |
|------|----------|
| `none` | No position IDs generated |
| `doc_reset` | Reset position to 0 at each document boundary in packed sequences |
| `continuous` | Continuous position IDs across packed documents |
Default is `doc_reset`, which ensures each document in a packed bin starts from position 0, preventing position encoding drift between unrelated documents.
## Gradient Accumulation Mechanics
Three cooperating layers enable gradient accumulation:
1. **`GradientState`** — tracks the micro-step counter. Fires `sync_gradients=True` every `grad_accum_steps` micro-batches. The counter is incremented at the **start** of `accumulate()`, before the forward pass.
2. **`executor._no_sync(model)`** — suppresses gradient synchronization on non-sync micro-steps:
- `NoneExecutor`: `nullcontext` (nothing to skip)
- `DDPExecutor`: `model.no_sync()` (PyTorch's built-in — skips all-reduce of gradient buckets)
- `FSDPExecutor`: `set_requires_gradient_sync(False, recurse=True)` on each `FSDPModule` (FSDP2's native mechanism)
3. **`AccumOptimizer` / `AccumScheduler`** — wrap the real optimizer/scheduler. `step()` and `zero_grad()` are gated on `sync_gradients` — they only forward to the inner optimizer when the sync flag is True.
The loss is divided by `grad_accum_steps` before `backward()`, so gradients sum to the correct mean across micro-steps. `consumed_samples` increments by `batch_per_device * world_size` every micro-batch.
### Effective batch size
$$ \text{Effective batch} = \text{nprocs} \times \text{batch\_per\_device} \times \text{grad\_accum\_steps} $$
### Total optimizer steps
```
samples_per_replica = ceil(dataset_len / nprocs)
batches_per_replica = ceil(samples_per_replica / batch_per_device)
total_steps = (batches_per_replica // grad_accum_steps) * n_epoch
```
This accounts for data-parallel sharding — each rank processes `1/nprocs` of the dataset.
> Document Update Time: 2026-08-02
+254
View File
@@ -0,0 +1,254 @@
# Getting Started
This guide walks you through installing AstrAI, downloading a model, running inference, preprocessing data, and launching your first training job.
## Contents
- [Prerequisites](#prerequisites)
- [1. Install](#1-install)
- [2. Download Model Weights](#2-download-model-weights)
- [3. Run Inference](#3-run-inference)
- [4. Preprocess Data](#4-preprocess-data)
- [5. Train](#5-train)
- [6. Evaluate](#6-evaluate)
- [7. Docker](#7-docker)
- [Next Steps](#next-steps)
## Prerequisites
- **Python 3.12+**
- **PyTorch 2.11.0** (the exact version pinned by AstrAI; CUDA 12.8 build recommended for GPU support)
- NVIDIA GPU with CUDA for training, `scripts/tools/generate.py`, generation evaluations, and demos. The HTTP server and direct-scoring evaluations can run on CPU where their CLI exposes a CPU device.
## 1. Install
```bash
git clone https://github.com/ViperEkura/AstrAI.git
cd AstrAI
# Basic install (pure PyTorch, no custom CUDA kernels)
pip install -e .
# With CUDA kernels (optional, for fused attention and rotary embedding)
# CSRC_KERNELS=true pip install -e . --no-build-isolation
# With dev dependencies (pytest, ruff)
# pip install -e ".[dev]"
```
> **CUDA kernels** are opt-in. They are not built by default. When built, they can be activated via `with attn_backend(ATTN_BACKEND.CUDA):` for accelerated decode/prefill, and the fused rotary embedding kernel is auto-dispatched when available. You can skip them for normal usage.
## 2. Download Model Weights
AstrAI uses HuggingFace-style model directories. Download the default 1B instruction-tuned model:
```bash
python scripts/demo/download.py
# → Downloads to params/
```
To use a different model:
```bash
python scripts/demo/download.py --repo-id <HF_REPO_ID> --local-dir ./my_model
```
The model directory contains:
- `config.json` — model architecture configuration
- `model.safetensors` — model weights
- `tokenizer.json` + `tokenizer_config.json` — tokenizer files (including chat template)
## 3. Run Inference
### Interactive Chat (Simplest)
```bash
python scripts/demo/stream_chat.py
# Type your message after >>, type !exit to quit
```
This starts a single-turn interactive prompt loop with streaming output. Each prompt is independent; conversation history is not retained.
### Start an HTTP Server
```bash
# Terminal 1: start server
python scripts/tools/server.py --param_path ./params --device cuda
# Terminal 2: query (OpenAI-compatible API)
curl -X POST http://localhost:8000/v1/chat/completions \
-H "Content-Type: application/json" \
-d '{"messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
```
The server also supports the Anthropic API at `/v1/messages`. See [Inference Guide](guides/inference.md) for full API documentation.
### Batch Generation from a File
Create an input JSONL file (one JSON object per line):
```json
{"question": "What is machine learning?"}
{"question": "Explain gradient descent."}
```
```bash
python scripts/tools/generate.py \
--param_path ./params \
--input_json_file input.jsonl \
--output_json_file output.jsonl
```
## 4. Preprocess Data
AstrAI uses a declarative JSON config to define the preprocessing pipeline. Create a config file for your training type:
### Pretraining (seq)
Input JSONL:
```json
{"text": "Artificial intelligence is..."}
```
Config (`pretrain.json`):
```json
{
"input": {
"sections": [{"field": "text", "action": "train"}]
},
"preprocessing": {"max_seq_len": 2048},
"output": {"storage_format": "bin"}
}
```
### SFT (Supervised Fine-Tuning)
Input JSONL:
```json
{"messages": [{"role": "user", "content": "Hi"}, {"role": "assistant", "content": "Hello!"}]}
```
Config (`sft.json`):
```json
{
"input": {
"sections": [{"field": "messages", "action": "$role", "template": true}]
},
"mask": {
"system": "mask",
"user": "mask",
"assistant": "train"
},
"mask_default": "mask",
"preprocessing": {"max_seq_len": 2048},
"output": {"storage_format": "bin", "dtype": {"loss_mask": "bool"}}
}
```
### Run Preprocessing
```bash
python scripts/tools/preprocess.py data/*.jsonl -o output/ -c pretrain.json
```
See [Preprocessing Guide](guides/preprocessing.md) for DPO/GRPO configs and all options.
## 5. Train
### Single GPU
```bash
python scripts/tools/train.py \
--train_type=seq \
--data_root_path=/path/to/dataset \
--param_path=./params \
--batch_per_device=4 \
--grad_accum_steps=8 \
--max_lr=1e-4 \
--window_size=2048 \
--ckpt_dir=./checkpoint \
--nprocs=1 \
--parallel_mode=none
```
### Multi-GPU (DDP)
```bash
export CUDA_VISIBLE_DEVICES=0,1,2,3
export NCCL_P2P_DISABLE=1
export NCCL_NET_GDR_LEVEL=0
python scripts/tools/train.py \
--train_type=seq \
--data_root_path=/path/to/dataset \
--param_path=./params \
--parallel_mode=ddp \
--nprocs=4 \
--batch_per_device=4 \
--grad_accum_steps=8 \
--max_lr=1e-4 \
--window_size=2048 \
--ckpt_dir=./checkpoint
```
### Training Types
| `--train_type` | Description | Data Keys |
|----------------|-------------|-----------|
| `seq` | Pre-training (next-token prediction) | `sequence` |
| `sft` | Supervised fine-tuning (masked loss) | `sequence`, `loss_mask` |
| `dpo` | Direct Preference Optimization | `chosen`, `rejected`, `*_mask` |
| `grpo` | Group Relative Policy Optimization | `prompts`, `responses`, `masks`, `rewards` |
See [Training Guide](guides/training.md) for loss formulas and strategies. See [Distributed Guide](guides/distributed.md) for DDP/FSDP details.
## 6. Evaluate
HumanEval and MMLU download their benchmark data through HuggingFace
`datasets`, which is not part of the base install:
```bash
pip install datasets
```
```bash
# HumanEval (code generation, auto-downloads dataset)
python scripts/eval/evaluate_humaneval.py --param_path ./params --num_samples 20
# MMLU (knowledge, auto-downloads dataset)
python scripts/eval/evaluate_mmlu.py --param_path ./params --n_shot 5
# Perplexity on custom data
python scripts/eval/evaluate_ppl.py --param_path ./params --input_path data.jsonl --output_dir ppl_results/
```
See [Evaluation Guide](guides/evaluation.md) for all benchmarks.
## 7. Docker
```bash
# Build
docker build -t astrai:latest .
# Run inference server with GPU
docker run --gpus all -p 8000:8000 astrai:latest \
python -m scripts.tools.server --port 8000 --device cuda
# Docker Compose (GPU)
docker compose up -d
```
## Next Steps
| Topic | Document |
|-------|----------|
| CLI parameters (train, server, generate, preprocess) | [CLI Reference](guides/params.md) |
| Preprocessing pipeline details | [Preprocessing Guide](guides/preprocessing.md) |
| Training loop, strategies, schedulers | [Training Guide](guides/training.md) |
| KV cache, continuous batching, HTTP API | [Inference Guide](guides/inference.md) |
| Evaluation benchmarks | [Evaluation Guide](guides/evaluation.md) |
| Multi-GPU DDP / FSDP | [Distributed Guide](guides/distributed.md) |
| System architecture | [Architecture](developer/architecture.md) |
| Data pipeline internals | [Data Flow](developer/dataflow.md) |
> Document Update Time: 2026-07-31
+263
View File
@@ -0,0 +1,263 @@
# Distributed Training
AstrAI supports three parallel modes: **single GPU** (`none`), **Data Parallel** (`ddp`), and **Fully Sharded Data Parallel** (`fsdp`). This guide covers when to use each, how to launch multi-GPU training, and how gradient accumulation works.
## Contents
- [Quick Start](#quick-start)
- [Parallel Modes](#parallel-modes)
- [Gradient Accumulation](#gradient-accumulation)
- [Process Launching](#process-launching)
- [NCCL Troubleshooting](#nccl-troubleshooting)
- [Checkpoint Saving](#checkpoint-saving)
- [Total Steps Calculation](#total-steps-calculation)
- [Real Examples](#real-examples)
- [CLI Parameters](#cli-parameters)
## Quick Start
### Single GPU
```bash
python scripts/tools/train.py \
--train_type=sft \
--param_path ./params \
--data_root_path ./dataset \
--parallel_mode=none \
--nprocs=1 \
--batch_per_device=4 \
--grad_accum_steps=8
```
### Multi-GPU DDP (4 GPUs)
```bash
export CUDA_VISIBLE_DEVICES=0,1,2,3
python scripts/tools/train.py \
--train_type=sft \
--param_path ./params \
--data_root_path ./dataset \
--parallel_mode=ddp \
--nprocs=4 \
--batch_per_device=4 \
--grad_accum_steps=8
```
### Multi-GPU FSDP (4 GPUs)
```bash
export CUDA_VISIBLE_DEVICES=0,1,2,3
python scripts/tools/train.py \
--train_type=sft \
--param_path ./params \
--data_root_path ./dataset \
--parallel_mode=fsdp \
--nprocs=4 \
--batch_per_device=4 \
--grad_accum_steps=8
```
> `--parallel_mode` defaults to `fsdp`. You can omit it for FSDP.
## Parallel Modes
| Mode | `--parallel_mode` | Param Layout | Memory | When to Use |
|------|-------------------|--------------|--------|-------------|
| Single GPU | `none` | Full, replicated | Highest | Small models, DPO/GRPO, debugging |
| DDP | `ddp` | Full, replicated | High | Most multi-GPU training |
| FSDP | `fsdp` | Sharded (DTensor) | Lowest | Large models that don't fit in single GPU |
### NoneExecutor
No wrapping. The model runs as-is on a single device. Gradient accumulation still works via `AccumOptimizer`/`AccumScheduler` (they gate `step()` on the sync counter). Checkpoint saving is a plain `state_dict()` call.
### DDPExecutor
Wraps the model with `torch.nn.parallel.DistributedDataParallel`. Each rank has a full copy of the model; gradients are all-reduced across ranks. Uses `gradient_as_bucket_view=True` and `broadcast_buffers=False` by default (hardcoded in `train.py`).
During gradient accumulation, non-sync micro-steps use `model.no_sync()` to skip gradient all-reduce. Only the final micro-step triggers the all-reduce.
### FSDPExecutor (FSDP2 / `fully_shard`)
Uses PyTorch's FSDP2 per-module API (`torch.distributed.fsdp.fully_shard`). Each model child (e.g., each `DecoderBlock`) is individually sharded — parameters become `DTensor`s distributed across ranks. No `FlatParameter`, original parameter names are preserved.
Key differences from DDP:
- **Lower memory**: parameters are sharded, not replicated.
- **Custom grad norm**: FSDP gradients are `DTensor`s, so `clip_grad_norm` computes the local norm, then all-reduces to get the global norm.
- **Collective checkpoint ops**: `unshard()` and `full_tensor()` are collective — all ranks must call them even though only rank-0 saves. The executor handles this via `dist.barrier()` in `checkpoint_context`.
- **Root skipped**: `fully_shard` is applied to direct children only (not the root model) due to an `ABC + Generic[T]` MRO incompatibility.
## Gradient Accumulation
Gradient accumulation lets you simulate a larger effective batch size by accumulating gradients over multiple micro-batches before calling `optimizer.step()`.
```
Effective batch = nprocs × batch_per_device × grad_accum_steps
```
Example: 4 GPUs × batch 4 × accum 8 = effective batch 256.
### How it works
Three cooperating layers:
1. **`GradientState`** — tracks the micro-step counter. Fires `sync_gradients=True` every `grad_accum_steps` micro-batches.
2. **`executor._no_sync(model)`** — suppresses gradient synchronization on non-sync micro-steps:
- `none`: `nullcontext` (nothing to skip)
- `ddp`: `model.no_sync()` (skips all-reduce)
- `fsdp`: `set_requires_gradient_sync(False)` on each `FSDPModule`
3. **`AccumOptimizer` / `AccumScheduler`** — gate `step()` and `zero_grad()` on `sync_gradients`, so the optimizer only fires on the last micro-step.
The loss is divided by `grad_accum_steps` before `backward()`, so gradients sum to the correct mean.
## Process Launching
AstrAI auto-detects the launch method:
| Detection | Strategy | Use Case |
|-----------|----------|----------|
| `torchelastic` / `torchrun` env vars | `TorchrunStrategy` | External orchestrator (`torchrun`, K8s) |
| `RANK` + `WORLD_SIZE` env vars | `TorchrunStrategy` | External launch |
| Neither | `LocalStrategy` | `python scripts/tools/train.py` (in-process spawn) |
### Local (default)
When you run `python scripts/tools/train.py --nprocs=4`, AstrAI uses `torch.multiprocessing.start_processes` to spawn 4 child processes. The parent process manages signal forwarding (SIGTERM/SIGINT) and waits for all children to finish.
### Torchrun
For multi-node or SLURM environments:
```bash
torchrun --nproc_per_node=4 scripts/tools/train.py \
--train_type=sft \
--parallel_mode=ddp \
--nprocs=4 \
--param_path ./params \
--data_root_path ./dataset \
--batch_per_device=4
```
When launched via `torchrun`, the launcher creates the worker processes. AstrAI reads `RANK`, `WORLD_SIZE`, and `LOCAL_RANK` from the environment and uses `TorchrunStrategy`; `--nprocs` does not control process creation in this mode.
The current training CLI still uses `--nprocs` when calculating scheduler `total_steps`. Set it to the global `WORLD_SIZE` so the step count reflects data-parallel sharding, including multi-node runs.
Raw Slurm variables such as `SLURM_PROCID`, `SLURM_NTASKS`, and `SLURM_LOCALID` are not recognized automatically. Launch through `torchrun`, or map the scheduler's variables to `RANK`, `WORLD_SIZE`, `LOCAL_RANK`, `MASTER_ADDR`, and `MASTER_PORT` before starting AstrAI. The same requirement applies to launchers that expose only OpenMPI-specific variables.
## NCCL Troubleshooting
The following variables are troubleshooting options for hardware or network configurations where NCCL hangs or fails. They are not general requirements and can reduce performance by disabling peer-to-peer or GPUDirect RDMA paths:
```bash
export NCCL_P2P_DISABLE=1
export NCCL_NET_GDR_LEVEL=0
```
Apply them only after confirming the relevant NCCL transport is the source of the failure. AstrAI does not set them in Python.
## Checkpoint Saving
Checkpoints are saved by **rank-0 only**. The flow:
1. `executor.checkpoint_context(model)` — wraps with `dist.barrier()` before and after (distributed only).
2. `executor.unwrap_model(model)` — gathers the full state dict:
- `none`: `model.state_dict()`
- `ddp`: `model.module.state_dict()`
- `fsdp`: `unshard()``full_tensor()``reshard()` (collective on all ranks, result kept only on rank-0)
3. Non-rank-0 ranks get `None` — the save is skipped.
4. Rank-0 writes `meta.json`, `config.json`, `model.safetensors`, and optional `{key}.pt` (optimizer/scheduler state).
> **FSDP note**: Even though only rank-0 saves, all ranks must participate in `unwrap_model` because `unshard()` and `full_tensor()` are collective operations. The barriers in `checkpoint_context` keep all ranks in lockstep.
## Total Steps Calculation
The scheduler's total step count accounts for data-parallel sharding:
```
samples_per_replica = ceil(dataset_len / nprocs)
batches_per_replica = ceil(samples_per_replica / batch_per_device)
total_steps = (batches_per_replica // grad_accum_steps) * n_epoch
```
This ensures the LR schedule is correctly scaled regardless of the number of GPUs.
## Real Examples
### Pretraining (seq, DDP, 4 GPUs)
```bash
export CUDA_VISIBLE_DEVICES=0,1,2,3
python scripts/tools/train.py \
--train_type=seq \
--param_path ./params \
--data_root_path ./dataset/cached \
--parallel_mode=ddp \
--nprocs=4 \
--n_epoch=1 \
--max_lr=2e-4 \
--schedule_type=wsd \
--warmup_ratio=0.02 \
--window_size=2048 \
--batch_per_device=4 \
--grad_accum_steps=32 \
--ckpt_interval=2000
# Effective batch = 4 × 4 × 32 = 512
```
### SFT (DDP, 4 GPUs)
```bash
python scripts/tools/train.py \
--train_type=sft \
--param_path ./AstrAI-V1-base \
--data_root_path ./dataset/cached_sft \
--parallel_mode=ddp \
--nprocs=4 \
--n_epoch=2 \
--max_lr=2e-5 \
--schedule_type=cosine \
--warmup_ratio=0.02 \
--min_rate=0.05 \
--window_size=2048 \
--batch_per_device=4 \
--grad_accum_steps=8
# Effective batch = 4 × 4 × 8 = 128
```
### DPO (Single GPU)
```bash
python scripts/tools/train.py \
--train_type=dpo \
--param_path ./checkpoint/epoch_1_step_6000 \
--data_root_path ./alpaca_dpo.jsonl \
--parallel_mode=none \
--nprocs=1 \
--max_lr=5e-6 \
--schedule_type=cosine \
--warmup_ratio=0.1 \
--min_rate=0.1 \
--window_size=1024 \
--batch_per_device=4 \
--grad_accum_steps=8 \
--dpo_beta=0.1 \
--max_grad_norm=50
```
## CLI Parameters
| Parameter | Default | Description |
|-----------|---------|-------------|
| `--nprocs` | 1 | Local process count for AstrAI's launcher; under `torchrun`, set it to global `WORLD_SIZE` for total-step calculation |
| `--parallel_mode` | `fsdp` | `none`, `ddp`, or `fsdp` |
| `--start_method` | `spawn` | Multiprocessing start method (`spawn`, `fork`, `forkserver`) |
| `--backend` | `nccl` | Distributed backend (`nccl`, `gloo`) |
| `--master_addr` | `localhost` | Master node address |
| `--master_port` | `29500` | Master node port |
| `--device_type` | `cuda` | Device type |
> `--tp_size` is accepted by the CLI but discarded before configuration. Tensor parallelism is not implemented, and there is no tensor-parallel module or model integration.
Full parameter reference: [CLI Reference](params.md). Training loop and strategies: [Training Guide](training.md).
> Document Update Time: 2026-08-02
+286
View File
@@ -0,0 +1,286 @@
# Evaluation
AstrAI provides 7 evaluation scripts in `scripts/eval/` covering code generation, knowledge QA, perplexity, summarization, data quality, instruction following, and weight analysis.
## Contents
- [Prerequisites](#prerequisites)
- [Overview](#overview)
- [HumanEval](#humaneval-code-generation)
- [MMLU](#mmlu-knowledge-qa)
- [Perplexity](#perplexity-ppl)
- [ROUGE](#rouge)
- [IFD](#ifd-instruction-following-difficulty)
- [IFEval](#ifeval-instruction-following)
- [Weight Analysis](#weight-analysis)
- [Tips](#tips)
## Prerequisites
HumanEval, MMLU, and IFEval import HuggingFace `datasets` to download their benchmark data. This package is not installed by AstrAI's base dependencies, so install it before running those scripts:
```bash
pip install datasets
```
The generation-based scripts require CUDA because they load the model on `cuda` with `bfloat16`. Direct-scoring and metric scripts support the devices shown below.
## Overview
| Script | Metric | Model Invocation | External Dataset |
|--------|--------|-------------------|-------------------|
| `evaluate_humaneval.py` | Code-gen pass@1/10/100 | `InferenceEngine.generate` | HF `openai/openai_humaneval` (auto-download) |
| `evaluate_mmlu.py` | MCQ accuracy (log-likelihood) | Direct `model()` forward | HF `cais/mmlu` (auto-download) |
| `evaluate_ppl.py` | Perplexity / token loss | Direct `model()` forward | User JSONL |
| `evaluate_rouge.py` | ROUGE-1/2/L | None (pure metric) | User JSONL |
| `evaluate_ifd.py` | Instruction-Following Difficulty | Direct `model()` forward | User JSONL |
| `evaluate_ifeval.py` | Instruction-following constraints | `InferenceEngine.generate` | HF `google/IFEval` (auto-download) |
| `analyze_weights.py` | SVD effective rank / weight stats | None (loads safetensors) | Checkpoint dir |
Two invocation patterns exist:
- **Generation benchmarks** (HumanEval, IFEval): use `InferenceEngine` to generate responses, then score them.
- **Scoring benchmarks** (MMLU, PPL, IFD): call `model()` directly under `torch.inference_mode()` for log-likelihood computation.
| Script | Device support |
|--------|----------------|
| HumanEval | CUDA for generation; `--test_only` can score existing completions without loading a model |
| IFEval | CUDA only |
| MMLU | CUDA or CPU via `--device`; auto-selects CUDA when available |
| PPL | CUDA or CPU via `--device`; auto-selects CUDA when available |
| IFD | CUDA or CPU via `--device`; auto-selects CUDA when available |
| ROUGE | CPU-only metric computation; no model is loaded |
| Weight analysis | CUDA by default; CPU supported via `--device cpu` |
---
## HumanEval (Code Generation)
Generates completions for 164 programming problems, executes them against hidden tests, and reports pass@k.
```bash
python scripts/eval/evaluate_humaneval.py \
--param_path ./params \
--num_samples 20 \
--batch_size 64 \
--max_tokens 512 \
--output results/humaneval.json
```
| Parameter | Default | Description |
|-----------|---------|-------------|
| `--param_path` | `./params` | Model directory |
| `--data_path` | `./humaneval/HumanEval.jsonl` | HumanEval JSONL (auto-downloaded if missing) |
| `--output` | None | Save results JSON (also writes `_completions.json`) |
| `--test_only` | None | Test an existing completions JSON (skip generation) |
| `--generate_only` | False | Only generate, skip execution/testing |
| `--num_samples` | 200 | Completions per problem (pass@k needs >= k) |
| `--max_tokens` | 512 | Max generation length |
| `--temperature` | 0.8 | Sampling temperature |
| `--top_p` | 0.95 | Nucleus sampling threshold |
| `--top_k` | 50 | Top-k sampling |
| `--batch_size` | 64 | Generation batch size |
| `--max_seq_len` | 4096 | KV cache sequence length |
| `--test_workers` | 8 | ProcessPoolExecutor workers for test execution |
| `--test_timeout` | 3.0 | Per-subprocess timeout (seconds) |
| `--problems` | None | Restrict to specific problem indices |
**Output**: stdout prints `pass@1`, `pass@10`, `pass@100`. With `--output`, writes per-problem results + `_summary` aggregate and a `_completions.json` file.
**Data**: Auto-downloads `openai/openai_humaneval` from HuggingFace on first run. Each problem has `task_id`, `entry_point`, `prompt`, `test`.
---
## MMLU (Knowledge QA)
57-subject multiple-choice accuracy via log-likelihood comparison. Supports n-shot few-shot prompting and option permutation.
```bash
python scripts/eval/evaluate_mmlu.py \
--param_path ./params \
--n_shot 5 \
--subjects abstract_algebra high_school_us_history \
--output results/mmlu.json
```
| Parameter | Default | Description |
|-----------|---------|-------------|
| `--param_path` | `./params` | Model directory |
| `--data_dir` | `./mmlu_data` | MMLU data directory (per-subject CSVs) |
| `--download` | False | Force re-download |
| `--n_shot` | 5 | Few-shot examples (0 = zero-shot) |
| `--subjects` | all 57 | Specific subjects to evaluate |
| `--output` | None | Output JSON path |
| `--split` | `test` | `test` or `val` |
| `--device` | auto | Device (`cuda` / `cpu`) |
| `--dtype` | auto | `bfloat16` on CUDA, `float32` on CPU |
| `--seed` | 0 | Seed for option permutation (0 = enabled, -1 = disabled) |
| `--batch_size` | 4 | Questions per batch; each question produces four choice rows |
**How it works**: For each question, builds a prompt with n-shot examples, then scores each choice (A/B/C/D) by computing the summed log-likelihood of the choice token given the context. The choice with the highest log-prob is the prediction.
**Output**: stdout prints per-subject accuracy and overall. With `--output`, writes per-subject `{accuracy, correct, total}` + `_overall` aggregate.
**Data**: Auto-downloads `cais/mmlu` from HuggingFace. Stored as per-subject CSVs in `<data_dir>/<split>/` and `<data_dir>/dev/` (for few-shot). `--subjects` accepts canonical MMLU names such as `abstract_algebra`, `college_computer_science`, `high_school_us_history`, and `world_religions`.
---
## Perplexity (PPL)
Token-level negative-log-likelihood and perplexity on arbitrary text data. Supports streaming mode (memory-efficient) and non-streaming mode (exact per-token stats).
```bash
python scripts/eval/evaluate_ppl.py \
--param_path ./params \
--input_path data.jsonl \
--output_dir ppl_results/ \
--batch_size 64 \
--max_length 2048
```
| Parameter | Default | Description |
|-----------|---------|-------------|
| `--param_path` | required | Model directory |
| `--input_path` | required | Input file, glob, or directory |
| `--output_dir` | required | Output directory for `summary.json` + token JSONL |
| `--text_key` | `text` | Key for the text field in input data |
| `--batch_size` | 64 | Batch size |
| `--max_length` | 2048 | Max sequence length (tokens) |
| `--token_level` | False | Store per-token log_probs + token-type analysis |
| `--max_samples` | None | Random subsample per file |
| `--device` | auto | Device |
| `--dtype` | auto | Torch dtype |
**Input**: JSONL or JSON files. Each item must have a field named by `--text_key` (default `text`). If `--input_path` is a directory, recursively collects `*.jsonl` and `*.json`.
**Output**: `summary.json` with per-file token count, mean loss, perplexity, and p50/p90/p95/p99 loss. Median loss is included only with `--token_level`; that mode also writes per-token JSONL with token IDs and log-probs.
---
## ROUGE
ROUGE-1/2/L (precision, recall, F1) for summarization. Self-contained implementation with no external dependencies.
```bash
python scripts/eval/evaluate_rouge.py \
--data_path predictions.jsonl \
--output results/rouge.json
```
| Parameter | Default | Description |
|-----------|---------|-------------|
| `--data_path` | required | JSONL with `reference`/`candidate` per line |
| `--output` | None | Output JSON path |
**Input**: JSONL, one object per line:
```json
{"reference": "Ground truth text", "candidate": "Model output text"}
```
**Output**: stdout prints `rouge-1`, `rouge-2`, `rouge-l` each as P/R/F1. With `--output`, writes JSON with `aggregate` and `per_item` scores.
Can also be imported as a library:
```python
from scripts.eval.evaluate_rouge import compute_rouge
scores = compute_rouge(reference, candidate)
```
---
## IFD (Instruction-Following Difficulty)
Data quality metric: `IFD = L_conditional / L_unconditional`. Measures how much harder it is to predict a response given its instruction vs. without it. Useful for filtering instruction-tuning data.
```bash
python scripts/eval/evaluate_ifd.py \
--param_path ./params \
--input_path sft_data.jsonl \
--output_dir ifd_results/ \
--format messages \
--batch_size 8
```
| Parameter | Default | Description |
|-----------|---------|-------------|
| `--param_path` | required | Model directory |
| `--input_path` | required | Input file, glob, or directory |
| `--output_dir` | required | Output directory |
| `--max_len` | 2048 | Max token length |
| `--format` | `plain` | `plain` (instruction/response fields) or `messages` (chat format) |
| `--instr_key` | `instruction` | Instruction field key (plain format) |
| `--resp_key` | `response` | Response field key (plain format) |
| `--batch_size` | 8 | Items per model-forward flush |
| `--device` | auto | Device |
| `--dtype` | auto | Torch dtype |
| `--sentinel_text` | `\n` | Prefix for unconditional pass (`""` → bos/pad fallback) |
| `--per_token` | False | Include per-token IFD breakdown |
| `--max_samples` | None | Random subsample per file |
**How it works**: Two forward passes per batch — (1) conditional: packed BFD sequence with context + response, (2) unconditional: response prefixed with a sentinel. IFD = mean_conditional_loss / mean_unconditional_loss. IFD > 1 means the instruction makes the response harder to predict (higher quality data).
**Output**: Per-file `<label>_ifd.jsonl` with IFD scores per item. `summary.json` aggregates per-file stats.
---
## IFEval (Instruction Following)
Google's IFEval benchmark: generates responses and verifies 27 types of constraints (keywords, format, length, case, punctuation, etc.).
```bash
python scripts/eval/evaluate_ifeval.py \
--param_path ./params \
--num_samples 1 \
--max_tokens 512 \
--output results/ifeval.json
```
| Parameter | Default | Description |
|-----------|---------|-------------|
| `--param_path` | `./params` | Model directory |
| `--data_path` | `./ifeval/input_data.jsonl` | IFEval JSONL (auto-downloaded if missing) |
| `--output` | None | Output JSON path |
| `--max_tokens` | 512 | Max generation tokens |
| `--temperature` | 0.1 | Sampling temperature (low for instruction-following) |
| `--top_p` | 0.95 | Top-p sampling |
| `--top_k` | 50 | Top-k sampling |
| `--num_samples` | 1 | Samples per problem (best-of-n scoring) |
| `--batch_size` | 64 | Inference batch size |
| `--max_seq_len` | 4096 | KV cache sequence length |
| `--limit` | None | Limit to first N problems (quick testing) |
| `--dump_responses` | None | Path to dump raw responses as JSONL |
**Output**: stdout prints overall accuracy + per-constraint-type accuracy table. With `--output`, writes per-problem results + `_summary`.
**Data**: Auto-downloads `google/IFEval` from HuggingFace. Each problem has `key`, `prompt`, `instruction_id_list`, `kwargs`.
---
## Weight Analysis
SVD-based effective rank and weight statistics for checkpoint diagnostics. Does not load the model graph or run any forward pass.
```bash
python scripts/eval/analyze_weights.py \
--ckpt_dir ./checkpoint/epoch_1_step_6000 \
--output results/weights.json
```
| Parameter | Default | Description |
|-----------|---------|-------------|
| `--ckpt_dir` | required | Checkpoint directory containing `model.safetensors` |
| `--compare` | None | Additional checkpoint dirs to compare |
| `--no_svd` | False | Skip SVD; show only weight stats (faster) |
| `--output` | None | Save results as JSON |
| `--device` | `cuda` | Device for SVD |
**Output**: SVD effective rank by component (ER@90/95/99%, entropic rank, condition number), per-layer effective rank grid, and weight value statistics (mean/std/min/max). Provides a utilization verdict (HIGH >0.85 / MODERATE >0.5 / LOW).
---
## Tips
- **Quick test**: Use `--limit` (IFEval) or `--problems` (HumanEval) to run on a small subset first.
- **Auto-download**: After installing `datasets`, HumanEval, MMLU, and IFEval auto-download their datasets on first run. The other scripts expect user-provided data.
- **Output formats**: `--output` writes a single JSON for most scripts. PPL and IFD write an `--output_dir` containing `summary.json` plus per-file artifacts.
- **CPU mode**: MMLU, PPL, and IFD support `--device cpu --dtype float32`; weight analysis supports `--device cpu`. HumanEval generation and IFEval are CUDA-only.
> Document Update Time: 2026-07-30
@@ -4,6 +4,7 @@
- [KV Cache](#kv-cache)
- [KVCache System](#kvcache-system)
- [Attention Backend](#attention-backend)
- [Continuous Batching](#continuous-batching)
- [Sampling](#sampling-strategy-pattern)
- [Protocol Handlers](#protocol-handlers-strategy-pattern)
@@ -23,30 +24,71 @@ RoPE is applied **before** KV cache write, not after — otherwise position enco
## KVCache System
Seven classes working together, with two concrete cache implementations:
### ContiguousCache (default)
Three-layer separation (SGLang-inspired): storage, index table, allocator.
```
ContiguousCache (simple contiguous per-slot cache)
├── ContiguousCacheView bundles k/v tensors + slot indices for attention layers
PagePool (top-level manager, orchestrates all layers)
├── KVStorage k_buffer / v_buffer [n_layers, size, n_kv_heads, head_dim]
├── ReqToTokenPool req_to_token [num_reqs, max_ctx_len] → physical token slot
├── Allocator bitmask-based page allocator + ref-count + LRU (paged mode only)
└── PrefixCache hash-based prefix matching (paged mode only)
```
Created by default when no cache is passed to `InferenceScheduler`. Each task occupies a fixed slot of `[max_seq_len, num_key_value_heads, head_dim]`. Simple and efficient for small-to-medium batch sizes.
`PagePool` supports two modes:
### PageCache (paged with prefix sharing)
- **Contiguous (default)**: pre-allocates `max_batch_size * max_seq_len` token slots. `req_to_token` is a trivial linear mapping (`slot = req_idx * max_seq_len + pos`). No dynamic allocation.
- **Paged** (`page_size=1` or `>1` with `n_tokens` set): shared token pool with on-demand allocation. Allocator + PrefixCache enable prefix sharing and LRU eviction.
`bind_tasks()` returns a `KVCache` dataclass — pure data, no methods:
```
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 (num_hidden_layers × n_pages × page_size × num_key_value_heads × head_dim)
── PageCacheView bundles Storage + page_table + total_len for attention layers
KVCache
├── k_buffer, v_buffer [n_layers, size, n_kv_heads, head_dim]
├── req_to_token [num_reqs, max_ctx_len]
── req_pool_indices [batch_size]
├── seq_lens [batch_size]
├── out_cache_loc [batch, seq_len] — write indices for this forward
── max_len int — max(seq_lens), avoids GPU sync in decode
└── kv_indptr [batch + 1] int32 — prefix sum of seq_lens, precomputed once per step
```
`isinstance(cache, KVCache)` checks dispatch to the correct view. Both implement the abstract `KVCache` interface used by `Executor` and `InferenceScheduler`.
Attention layers do raw buffer indexing: `k_buffer[layer_id, out_cache_loc] = k` to write, `k_buffer[layer_id, indices]` to gather.
## Attention Backend
Attention computation (cache I/O + SDPA/kernel dispatch) is decoupled from the model via `AttentionBackend` ABC:
```
AttentionBackend (ABC)
├── TorchNativeBackend SDPA + indirect KV cache gather (default)
└── CudaBackend CUDA kernel dispatch (attn_paged_decode, attn_prefill)
```
Select 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` decode path: writes K/V to cache, then calls `attn_paged_decode` with `page_size=1` — the `req_to_token` table serves directly as the page table, each token slot is a single-token "page". No explicit K/V gather needed.
`CudaBackend` prefill path: writes K/V, gathers full-sequence K/V via indirect indexing (same as `TorchNativeBackend`), then calls `attn_prefill`.
Fallback: `CudaBackend` delegates to `TorchNativeBackend` when a CUDA kernel is not available.
### Rotary Embedding Backend
Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/rotary_backend.py`, which auto-dispatches:
- **CUDA kernel** (`rotary_emb.cu`): fused cos/sin lookup + rotation in a single kernel, used when the kernel is available, input is on CUDA, and `torch.is_grad_enabled()` is `False` (inference mode)
- **Torch fallback**: complex multiply path (`torch.view_as_complex``torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd backward) or when the CUDA kernel is not available
`RotaryEmbedding` stores a complex `freqs_cis` buffer and returns a tensor
from `forward()`. Both attention backends share the same rotary dispatch — it
is backend-agnostic.
## Continuous Batching
@@ -142,17 +184,57 @@ curl -X POST http://localhost:8000/v1/messages \
-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`.
Supports `stop_sequences` and streaming via `event: content_block_delta`. Anthropic streams also end with the shared `data: [DONE]` sentinel after `event: message_stop`.
### GenerationRequest Parameters
### Request Parameters
The HTTP protocols and direct engine API have distinct request models and defaults.
**OpenAI** (`ChatCompletionRequest`):
| Param | Type | Default | Description |
|-------|------|---------|-------------|
| `model` | str | `"astrai"` | Model name returned in responses |
| `messages` | List[dict] | required | Chat messages (role, content) |
| `top_k` | int | 50 | Top-k count |
| `temperature` | Optional[float] | 1.0 | Sampling temperature (0.0-2.0) |
| `top_p` | Optional[float] | 1.0 | Nucleus threshold (0.0-1.0) |
| `top_k` | Optional[int] | 50 | Top-k count |
| `max_tokens` | Optional[int] | 2048 | Max generation length |
| `stream` | Optional[bool] | False | Stream output |
| `stop` | Optional[Union[str, List[str]]] | None | Stop sequences |
| `n` | Optional[int] | 1 | Number of choices requested |
| `presence_penalty` | Optional[float] | 0.0 | Presence penalty (-2.0 to 2.0) |
| `frequency_penalty` | Optional[float] | 0.0 | Frequency penalty (-2.0 to 2.0) |
| `logit_bias` | Optional[Dict[int, float]] | None | Per-token logit bias |
| `user` | Optional[str] | None | End-user identifier |
| `tools` | Optional[List[ToolDef]] | None | Tool definitions for function calling |
| `tool_choice` | Optional[Union[str, Dict[str, Any]]] | `"auto"` | Tool selection mode or explicit tool choice |
**Anthropic** (`MessagesRequest`):
| Param | Type | Default | Description |
|-------|------|---------|-------------|
| `model` | str | `"astrai"` | Model name returned in responses |
| `messages` | List[AnthropicMessage] | required | User/assistant messages |
| `system` | Optional[str] | None | System prompt |
| `max_tokens` | int | 1024 | Max generation length |
| `temperature` | Optional[float] | 1.0 | Sampling temperature (0.0-2.0) |
| `top_p` | Optional[float] | 1.0 | Nucleus threshold (0.0-1.0) |
| `top_k` | Optional[int] | 50 | Top-k count |
| `stream` | Optional[bool] | False | Stream output |
| `stop_sequences` | Optional[List[str]] | None | Stop sequences |
**Engine** (`GenerationRequest`):
| Param | Type | Default | Description |
|-------|------|---------|-------------|
| `messages` | List[Dict[str, str]] | required | Messages to format before generation |
| `top_k` | int | 50 | Top-k count; 0 disables filtering |
| `top_p` | float | 1.0 | Nucleus threshold |
| `temperature` | float | 1.0 | Sampling temperature (> 0.0) |
| `temperature` | float | 1.0 | Sampling temperature; 0 enables greedy decoding |
| `max_tokens` | Optional[int] | None | Max generation length |
| `frequency_penalty` | float | 0.0 | Frequency penalty (-2.0 to 2.0) |
| `rep_window` | int | 64 | Recent-token window used by the frequency penalty |
| `stream` | bool | False | Stream output |
### SSE Streaming Format
@@ -195,6 +277,8 @@ data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":
event: message_stop
data: {"type":"message_stop"}
data: [DONE]
```
### Error Responses
@@ -249,4 +333,4 @@ async for token in engine.generate_async("Hello", ...): # -> AsyncGenerator[s
print(token)
```
> Document Update Time: 2026-07-09
> Document Update Time: 2026-07-31
+71 -21
View File
@@ -13,9 +13,11 @@
| Parameter | Description | Default |
|-----------|-------------|---------|
| `--config`, `-c` | YAML config file; explicit CLI options override YAML values | None |
| `--train_type` | Training type (`seq`, `sft`, `dpo`, `grpo`, `online_grpo`, `online_dpo`) | required |
| `--data_root_path` | Dataset root directory | required |
| `--param_path` | Model parameters or checkpoint path | required |
| `--resume` | Resume training from `--param_path` | False |
| `--n_epoch` | Total training epochs | 1 |
| `--batch_per_device` | Batch size per device | 1 |
| `--grad_accum_steps` | Gradient accumulation steps between optimizer steps | 1 |
@@ -26,20 +28,53 @@
|-----------|-------------|---------|
| `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 |
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
| `--max_grad_norm` | Maximum gradient norm for clipping (None disables) | 1.0 |
| `--max_grad_norm` | Maximum gradient norm for clipping; the current CLI requires a positive number | 1.0 |
### Optimizer (MuonMix)
### Optimizer
Combined optimizer: matrix parameters via **Muon**, non-matrix via **AdamW** (`fused=True`).
The default `muon_adamw` optimizer sends matrix parameters through **Muon** and
non-matrix parameters through **AdamW** (`fused=True`).
| Parameter | Description | Default |
|-----------|-------------|---------|
| `--weight_decay` | Weight decay (applied to Muon matrix params; non-matrix use 0) | 0.1 |
| `--optimizer` | Built-in optimizer (`muon_adamw`, `nora_nadamw`, `mano_adamw`) | `muon_adamw` |
| `--weight_decay` | Weight decay for optimizer parameter groups that are eligible for decay | 0.1 |
| `--muon_momentum` | Muon momentum factor | 0.95 |
| `--muon_nesterov` | Enable Nesterov momentum for Muon | True |
| `--muon_nesterov`, `--no-muon_nesterov` | Enable or disable Nesterov momentum for Muon | enabled |
| `--muon_ns_steps` | Newton-Schulz iteration steps for Muon | 5 |
| `--muon_adjust_lr` | Muon LR adjustment strategy (`original`, `match_rms_adamw`) | `match_rms_adamw` |
`nora_nadamw` routes internal `Linear.weight` matrices to **Nora** and
embeddings, the LM head, norms, biases, LoRA factors, and fallback parameters to
**NAdamW**. Parameters are classified by module role and identity, so tied
embedding/head weights occur in exactly one group. Nora requires complete rows
under DTensor sharding and rejects layouts sharded along the last dimension.
| Parameter | Description | Default |
|-----------|-------------|---------|
| `--nora_lr` | Nora learning rate | 5e-3 |
| `--nora_beta` | Nora momentum-buffer EMA factor | 0.95 |
| `--nora_momentum` | Nora Nesterov interpolation factor | 0.95 |
| `--nora_weight_decay` | Nora matrix weight decay | 0.0 |
`mano_adamw` routes internal `Linear.weight` matrices to **Mano** (manifold
normalized optimizer) and the remaining parameters to **AdamW**. Mano projects
the momentum onto the tangent space of the Oblique manifold and normalizes it,
alternating the projection axis (row/column) each step — replacing Muon's
Newton-Schulz iteration with a cheaper normalization.
| Parameter | Description | Default |
|-----------|-------------|---------|
| `--mano_momentum` | Accepted by the CLI but currently ignored by optimizer construction | 0.95 |
| `--mano_nesterov`, `--no-mano_nesterov` | Accepted by the CLI but currently ignored by optimizer construction | enabled |
The two Mano-specific flags are reserved for future wiring; do not rely on them
to change optimizer behavior in the current release.
Optimizer identity and hyperparameters are saved in checkpoint metadata. Optimizer
states are intentionally not interchangeable: resume older MuonAdamW checkpoints
with `--optimizer=muon_adamw`.
### Data Loading
| Parameter | Description | Default |
@@ -48,7 +83,7 @@ Combined optimizer: matrix parameters via **Muon**, non-matrix via **AdamW** (`f
| `--stride` | Stride for sliding window over sequences | None |
| `--random_seed` | Random seed for reproducibility | 3407 |
| `--num_workers` | DataLoader worker processes | 4 |
| `--no_pin_memory` | Disable pin_memory (enabled by default) | (flag) |
| `--pin_memory`, `--no-pin_memory` | Enable or disable DataLoader pinned memory | enabled |
### Checkpoint & Resume
@@ -70,41 +105,51 @@ Combined optimizer: matrix parameters via **Muon**, non-matrix via **AdamW** (`f
| Parameter | Description | Default |
|-----------|-------------|---------|
| `--log_dir` | Directory for metric logs | checkpoint/logs |
| `--metrics` | Metrics to log (e.g. --metrics loss lr val_loss) | ["loss", "lr", "grad_norm"] |
| `--metrics` | Repeatable metric option (for example, `--metrics loss --metrics lr --metrics val_loss`) | `loss`, `lr`, `grad_norm`, `grad_snr` |
### Gradient Checkpointing
| Parameter | Description | Default |
|-----------|-------------|---------|
| `--gradient_checkpointing` | Enable activation checkpointing for DecoderBlock modules | False |
| `--gradient_checkpointing`, `--no-gradient_checkpointing` | Enable or disable activation checkpointing for DecoderBlock modules | disabled |
### Miscellaneous
| Parameter | Description | Default |
|-----------|-------------|---------|
| `--compile` | Enable `torch.compile` with mode `default`, `reduce-overhead`, or `max-autotune`; omit to disable | None |
| `--dry-run` | Validate the merged configuration and print the training plan without training | False |
### Distributed Training
| Parameter | Description | Default |
|-----------|-------------|---------|
| `--nprocs` | Number of GPUs / processes | 1 |
| `--parallel_mode` | Parallel strategy (`none`, `ddp`, or `fsdp`) | none |
| `--parallel_mode` | Parallel strategy (`none`, `ddp`, `fsdp`) | fsdp |
| `--device_type` | Device type | cuda |
| `--start_method` | Multiprocessing start method (`spawn`, `fork`, `forkserver`) | spawn |
| `--backend` | Distributed training backend | nccl |
| `--master_addr` | Master node address | localhost |
| `--master_port` | Master node port | 29500 |
| `--tp_size` | Reserved tensor-parallel size; accepted but currently ignored | None |
### Strategy-specific
| Parameter | Description | Default | Used by |
|-----------|-------------|---------|---------|
| `--dpo_beta` | DPO beta value | 0.1 | `dpo` |
| `--dpo_beta` | DPO beta value | 0.1 | `dpo`, `online_dpo` |
| `--label_smoothing` | Label smoothing for cross-entropy loss | 0.0 | `seq`, `sft` |
| `--group_size` | GRPO group size | 4 | `grpo` |
| `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo` |
| `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 | `grpo` |
| `--group_size` | GRPO/rollout group size | 4 | `grpo`, `online_grpo`, `online_dpo` |
| `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo`, `online_grpo` |
| `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 | `grpo`, `online_grpo` |
| `--neftune_alpha` | NEFTune noise alpha (0=disabled, typical: 5.0) | 0.0 | `sft` |
### Online Rollout
These options apply to `online_grpo` and `online_dpo`. Online strategies require
`online_grpo` and `online_dpo` are factory aliases for the existing `grpo` and
`dpo` strategy classes; online behavior is enabled by rollout components rather
than separate strategy subclasses. These options apply to the online aliases.
Online strategies require
a `BaseRewardModel` factory in `TrainConfig`; `train.py` does not currently
provide a command-line option for configuring one.
@@ -121,7 +166,7 @@ provide a command-line option for configuring one.
| Parameter | Description | Default |
|-----------|-------------|---------|
| `--schedule_type` | LR scheduler type (`cosine`, `sgdr`, `wsd`) | cosine |
| `--min_rate` | Minimum LR as fraction of base LR | None (scheduler default: 0.05 for cosine/SGDR, 0.0 for WSD) |
| `--min_rate` | Minimum LR as fraction of base LR | None (all current schedulers use their effective default of 0.01) |
| `--cycle_length` | SGDR first cycle length in steps | None (total_steps - warmup_steps) |
| `--t_mult` | SGDR cycle length multiplier per restart | 2 |
| `--stable_steps` | WSD stable plateau steps | None (80% of post-warmup steps) |
@@ -164,6 +209,7 @@ nohup python scripts/tools/train.py \
| `--device` | str | `cuda` | Device to load model on |
| `--dtype` | str | `bfloat16` | Model weights dtype (`bfloat16`, `float16`, `float32`) |
| `--max_batch_size` | int | `16` | Maximum batch size for continuous batching |
| `--max_seq_len` | int | model config `max_position_embeddings` | Maximum sequence length (KV cache size + prompt truncation) |
| `--reload` | flag | `False` | Enable auto-reload for development |
Usage:
@@ -182,11 +228,14 @@ See [Inference Guide](inference.md) for HTTP API documentation.
| `--output_json_file` | str | required | Path to the output JSONL file |
| `--question_key` | str | `question` | Key for the question in input JSON |
| `--response_key` | str | `response` | Key for the response in output JSON |
| `--temperature` | float | `0.60` | Sampling temperature |
| `--top_k` | int | `30` | Top-k filtering |
| `--temperature` | float | `0.8` | Sampling temperature |
| `--top_k` | int | `50` | Top-k filtering |
| `--top_p` | float | `0.95` | Nucleus sampling threshold |
| `--batch_size` | int | `1` | Batch size for generation |
| `--max_tokens` | int | model config `max_position_embeddings` | Maximum tokens to generate |
| `--num_samples` | int | `1` | Responses per prompt |
| `--max_seq_len` | int | `2048` | KV cache sequence length |
| `--frequency_penalty` | float | `0.0` | Frequency penalty |
| `--rep_window` | int | `64` | Window size for frequency penalty |
Usage:
```bash
@@ -200,14 +249,15 @@ python scripts/tools/generate.py \
| Parameter | Type | Default | Description |
|-----------|------|---------|-------------|
| `input_files` | path(s) | required | Input JSONL file(s), supports glob (`data/*.jsonl`) |
| `input_files` | path(s) | required | One or more existing `.jsonl` or `.json` paths. Wildcards work only when expanded by the invoking shell; the CLI does not expand globs itself. |
| `--output_dir`, `-o` | path | required | Output directory for processed data |
| `--config`, `-c` | path | required | Preprocessing pipeline config (JSON) |
| `--tokenizer_path` | str | `params` | Path to tokenizer directory |
| `--batch_size` | int | config value (`256` by default) | Override records processed per batch; must be at least 1 |
Usage:
```bash
python scripts/tools/preprocess.py data/*.jsonl -o output/ -c sft.json
python scripts/tools/preprocess.py data/part-000.jsonl data/part-001.jsonl -o output/ -c sft.json
```
See [Preprocessing Guide](preprocessing.md) for config file format and examples.

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