56 Commits
Author SHA1 Message Date
ViperEkura b99485f462 chore: bump version to 1.3.11 2026-07-27 01:23:13 +08:00
ViperEkura 20041d7aa9 perf: extend MMA decode to arbitrary GQA ratio, add launch bounds, vectorize combine
- Multi-pass MMA: encode pass in grid blockIdx.x, compute q_head0/G in-kernel
- Fixes crash for G>32 (previously block(32,G) exceeded 1024 threads)
- Fixes alloc_split_partials using uninitialized num_splits (MAX_SPLITS=32)
- __launch_bounds__ on all MMA and prefill kernels for better register allocation
- 4x vectorized combine kernel (4 head_dim per thread)
- uint4 vectorized K loads in scalar decode kernels
- cp.async .L2::128B cache hint for K/V tile streaming
- Extract warp_reduce_sum, bf16, MAX_SPLITS to attn_warp_utils.cuh
2026-07-27 00:35:34 +08:00
ViperEkura 59248032dc chore: fix ruff lint warnings and signal handling edge cases
- Fix pre-existing ruff lint warnings (F401, F541, F841, E741)
- Exclude .md/.json/.yml from ruff format check
- Unblock SIGTERM/SIGINT via pthread_sigmask in early signal handler
- Do not restore SIG_DFL on unregister to prevent pending signal kills
2026-07-25 21:08:30 +08:00
ViperEkura ceadc34ea9 feat: auto-checkpoint on SIGTERM/SIGINT with DDP support
- Register SIGTERM/SIGINT handlers in training loop, set stop flag on signal
- Check stop_requested at each epoch/batch boundary, break and call on_error to save checkpoint
- LocalStrategy parent forwards signal to child processes via terminate(), waits up to 600s for graceful exit
- TrainContext gains threading.Event-based stop_requested/request_stop
- Tests verify SIGTERM/SIGINT trigger checkpoint save with exit code 0, works on both CPU and GPU
2026-07-25 20:40:54 +08:00
ViperEkura 8ab5631446 fix: correct online rollout lifecycle 2026-07-23 19:01:37 +08:00
ViperEkura 99b5d2b2da perf: batch tokenizer preprocessing 2026-07-23 18:42:19 +08:00
ViperEkura 021e6f3788 style: apply ruff formatting to FSDP2 changes 2026-07-23 16:30:10 +08:00
ViperEkura 4e38183e86 fix: make FSDP2 executor work with ABC+Generic model hierarchy
- Wrap each child module individually, skip root (CPython layout
  conflict between ABC+Generic and FSDP2 __class__ assignment)
- Remove manual unshard in clip_grad_norm (DTensor compatible)
- Fix _no_sync to iterate modules() instead of checking root
- Add reshard after unwrap_model
- Guard __init_subclass__ type resolution against dynamic subclasses
- Add fsdp2 to --parallel_mode CLI choices
2026-07-23 16:11:02 +08:00
ViperEkura 4eeb23e2b3 fix: use copy-on-write mmap mode to silence non-writable tensor warning 2026-07-22 17:37:41 +08:00
ViperEkura ef8783b7e3 fix: separate attn_mask and loss_mask in get_logprobs, compose causal masking in strategy
- add loss_mask parameter to get_logprobs to decouple attention from loss masking
- DPO/GRPO strategies compose key-padding + causal mask before model forward
- prevents prompt tokens from being masked out of attention and missing causal masking
2026-07-21 23:47:28 +08:00
ViperEkura 60d7ee614a fix: improve attention kernel numerical stability and test precision checks
- use fmaf() for V-accumulation in scalar decode paths to reduce rounding
- delay scale multiplication to after dot-product in scalar prefill
- unify __expf/expf across MMA and scalar paths for consistent numerics
- harmonize divide-by-zero guards to 1e-20f
- add both absolute and relative error checks in standalone tests (atol=0.01, rtol=0.01)
2026-07-21 23:05:17 +08:00
ViperEkura f7a16efc9d refactor: extract shared dispatcher header, unify MMA/scalar dispatch format
- Merge 3 duplicated dispatch blocks into single attn_dispatchers.cuh
- Merge compute_num_splits from attn_utils.cuh into dispatcher header
- All dim3 grid/block declarations and <<<>>> launches are single-line
- Production .cu files (35-42 loc) only handle torch wrapping + pybind11
- Test files include dispatcher header directly, removing all #ifndef ASTRAI_NO_MMA duplication
2026-07-21 22:21:39 +08:00
ViperEkura a01e8bbe98 refactor: adopt FA2-style KernelTraits + compile-time causal/mask dispatch
- Introduce KernelTraits<HEAD_DIM, BC, WARPS, STAGES> compile-time config bundle, replacing scattered <KD, NC8, KT2, ...> template params
- Template all MMA and scalar kernels on IsCausal/HasMask bools to eliminate inner-loop runtime branches
- Dispatch to 4-path IsCausal/HasMask kernel variants at entry points based on p.causal_offset and p.use_mask
- Update standalone test files with new kernel signatures, add causal test cases
- Fix duplicate using bf16 in MMA kernels that include attn_mma_utils.cuh
2026-07-21 21:52:46 +08:00
ViperEkura ccf728a1b7 perf: eliminate GPU syncs in contiguous cache write/gather hot paths
- Replace .tolist() calls with _total_len in gather(); move _slot_len updates from per-layer write to once-per-step bind_tasks

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

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

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

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-20 16:11:48 +08:00
ViperEkura 2c50b3cf37 ci: preserve both release wheel artifacts 2026-07-20 15:33:43 +08:00
ViperEkura eee7f54789 docs: sync training and architecture guides 2026-07-20 15:23:30 +08:00
ViperEkura 06eeeead79 refactor: map instruction/input/output to chat roles
- RolloutGenerator._instruction_to_messages builds system/user/assistant list (instruction->system, input->user, output->assistant), replacing single-user-turn concatenation
- Remove _iter_samples helper; _prepare_prompts zips parallel list-of-strings fields directly per the collate_fn contract
- Tests adopt a system-aware chat template and pin the three-field role mapping
- Drop unused imports caught by ruff F401 (torch.Tensor in scheduler.py, iter_raw_records in pipeline.py, Tuple in evaluate_rouge.py)
2026-07-20 13:55:25 +08:00
ViperEkura e8ff7f5321 fix: use batch_per_device for rollout scheduler batch sizing
- train_context.py referenced non-existent cfg.batch_size, replaced with cfg.batch_per_device
- default group_size lowered from 8 to 1: without a group concept (DPO), scheduler batch equals batch_per_device; rollout-based DPO can opt in via extra_kwargs['group_size']>=2
- inline expressions (rollout_batch_size, max_seq_len) extracted for readability
- add tests/trainer/test_online_e2e.py: end-to-end online_dpo via Trainer.train, exercising KV-cache-backed rollout path
2026-07-20 13:32:04 +08:00
ViperEkura a6e1f26cd4 refactor: simplify sample return_logprobs path
- SamplingPipeline.sample gains return_logprobs; both greedy and multinomial paths now share a single log_softmax+gather instead of duplicating the sampling logic
- module-level sample() becomes a thin forwarder instead of re-implementing the three-branch logic
- eliminates ~10 lines of duplicated softmax/gather code; no caller-facing API change
2026-07-20 13:16:18 +08:00
ViperEkura 95c43368ae refactor: unify rollout onto inference engine KV-cache path
- RolloutGenerator now delegates prefill/decode to InferenceScheduler.run_batch (sync API, no background thread), sharing one KV-cache code path with the inference server and eliminating O(n^2) recompute in rollout
- Add sample(return_logprobs=) and Executor.execute_decode(return_logprobs=) to expose behaviour-policy log-probs through the engine; Task gains output_logprobs
- RolloutResult now subclasses RawRollout (adds rewards only), removing duplicated fields
- RolloutRunner.__call__ returns (result, is_fresh) instead of relying on object identity, removing the fragile refresh-detection contract
- Remove O(n^2) generate_responses helper and dead code (_tokenize_prompts, unused old_model arg)
- train_context.py wires InferenceScheduler directly instead of hand-rolling SamplingPipeline
- Tests: +11 covering return_logprobs, run_batch, and KV-cache-backed rollout semantics; 404 pass
2026-07-20 12:52:20 +08:00
ViperEkura 754624acf0 feat: add online rollout framework for RL strategies
- RolloutRunner: generate + score responses with cached re-rollout trigger
- BaseStrategy.__call__ switches online/offline via runner injection
- GRPO/DPO implement prepare_from_rollout; aliases online_grpo/online_dpo
- TrainConfig + train.py add rollout params and CLI flags
- Tests cover generate_responses, RolloutRunner cache, shared __call__
2026-07-20 03:49:56 +08:00
ViperEkura 0b6a17330f feat: add FSDP2Executor using torch.distributed.fsdp.fully_shard API
- New FSDP2Executor registers as 'fsdp2' in ExecutorFactory, using per-module fully_shard() instead of FSDP1 FlatParameter wrapper
- FSDP2 preserves original Parameter objects as DTensors, eliminating use_orig_params=True hack
- FSDP2Executor implements _no_sync via set_requires_gradient_sync, clip_grad_norm via unshard, unwrap_model via DTensor.full_tensor
- Drop **_extra/**_ddp_only_kwargs fallbacks in BaseExecutor/FSDPExecutor/FSDP2Executor, replaced by parallel_mode-aware executor_kwargs dispatch in train.py (ddp-only kwargs only passed when parallel_mode=ddp)
- Export FSDP2Executor in astrai.parallel.__init__
2026-07-20 01:46:25 +08:00
ViperEkura 74b9308883 refactor: pass model_fn/optimizer_fn to executor.prepare
- BaseExecutor.prepare now takes factories and instantiates model via model_fn(), runs before_wrap hook, wraps DDP/FSDP, then builds optimizer/scheduler on the wrapped model
- optimizer/scheduler creation moved into executor.prepare, eliminating the old 'create-then-wrap' hack reliance on use_orig_params=True
- FSDPExecutor/BaseExecutor accept **_extra kwargs to tolerate DDP-only keys (broadcast_buffers, gradient_as_bucket_view) being forwarded via executor_kwargs
- dataloader builds stay external; executor only handles model/optimizer/scheduler
- train_context.py rewritten to load checkpoint state_dict before prepare via a before_wrap closure
2026-07-20 01:32:05 +08:00
ViperEkura e5f9b1a3a9 fix: default max_grad_norm to 1.0 and drop None branch 2026-07-20 01:08:13 +08:00
ViperEkura 31d33ccdf0 chore: bump to 1.3.10 2026-07-19 16:40:27 +08:00
ViperEkura 88ec786e39 fix: memmap mode=r, tool parser json.loads, greedy decode 2026-07-19 16:38:28 +08:00
ViperEkura 663ef900fc refactor: move sample-id indexing from dataset to store
- Store owns window_size/stride and __getitem__/__len__/sample_window
- Dataset classes become thin delegators binding a Store to a train-type key mapping
- Drop BaseDataset.get_index and the RecordDataset中间类 (window死代码)
- DatasetFactory forces window_size=0 for record datasets so record semantics never get window-tainted
- token_count/num_records split the legacy len() semantics (raw stream length vs record count)
- Update tests to the new .store/.token_count API and window/record mode switching
2026-07-19 16:02:50 +08:00
ViperEkura 7d478a54db docs: update HF org from ViperEk to ViperEkura
- Replace 4 HF links in README.md and README-zh-CN.md to point to ViperEkura
- Update download.py default repo to AstrAI-V1-instruct under ViperEkura
2026-07-19 14:49:58 +08:00
ViperEkura f3eaaef842 refactor: remove redundant strategy/executor code
- Drop BaseStrategy.model_fn (stored but never read)
- Drop model_fn= passed to StrategyFactory.create in train_context
- Simplify FSDPExecutor.clip_grad_norm None branch to delegate to super()
- Remove DDPExecutor._gather_state_dict override (identical to base)
2026-07-19 12:45:58 +08:00
ViperEkura d655b65027 docs: sync architecture/dataflow/training/params with code
- dataflow.md: update DatasetFactory.load signature, stream vs record access, Store._offsets
- architecture.md: add tokenizer to Pipeline, TokenizeTransform class, RecordDataset, Streamable/Recordable mixins, fix GRPOStrategy (old_model/sync_old_model)
- training.md: DPO reduction="sum", GRPO rho_t uses pi_old, gradient_clipping always registered
- params.md: --max_grad_norm default None
2026-07-19 12:33:35 +08:00
ViperEkura 31c22dc043 refactor: deduplicate preprocessing kernel and BFD packing
- Extract shared core (mask building, primary-id extraction, tensorisation, position-id generation) to astrai/preprocessing/core.py; Pipeline and TokenizeTransform both consume it, eliminating ~60% duplicated logic
- Promote BFD _plan to module-level plan_bfd(lengths, max_len) returning pure index bins; BFDPacking.apply and evaluate_ifd._pack_bins both call it, removing the second BFD implementation
- Split Pipeline._flush (49 lines) into _inject_doc_reset_position_ids + _inject_continuous_position_ids + _to_tensors; split Pipeline.run by delegating record iteration to core.iter_raw_records
- Remove dead no-op pop/塞回 in Pipeline.run (L110-111)
2026-07-19 12:27:56 +08:00
ViperEkura 17127f8b3c fix: make tokenizer picklable for spawn multiprocessing
- ChatTemplate: defer Jinja2 compilation to cached_property, exclude compiled template from __getstate__ (its dynamic root function has __module__=None and falls back to __main__, breaking pickle)
- AutoTokenizer: bypass __getattr__ for underscore-prefixed attrs to prevent infinite recursion during unpickle when __dict__ is empty
2026-07-19 11:59:55 +08:00
ViperEkura d7695b40e3 feat: make max_grad_norm optional (None disables clipping)
- TrainConfig.max_grad_norm defaults to None
- executor.clip_grad_norm returns grad norm without clipping when None
- train.py --max_grad_norm defaults to None
2026-07-19 00:08:18 +08:00
ViperEkura fc62890e70 fix: apply chat template in DPO tokenization
- dpo_tokenize now uses tokenizer.apply_chat_template to match SFT format
- Prompt rendered with add_generation_prompt=True
- Chosen/rejected appended as assistant turn
- Remove leftover dead code from _extract_text
- Update tests to mock apply_chat_template
2026-07-19 00:00:51 +08:00
ViperEkura f433672140 fix: use sum reduction for DPO sequence logprob
- DPO requires sequence-level sum of token logprobs, not per-token mean
- mean reduction made beta*ratio_diff ~0.03 (near-zero gradient)
- loss stalled at 0.6931 because logsigmoid(0.03) has vanishing grad
- sum gives beta*ratio_diff ~10 with meaningful gradients
2026-07-18 23:48:35 +08:00
ViperEkura 7e1e5b6e6a refactor: DatasetFactory.load accepts pre-built store instance
- load(store=...) binds directly, skipping format detection/processor
- load_path now optional when store is given
- Remove redundant from_store (merged into load)
- Caller can fully control Store construction + processor setup
2026-07-18 23:23:51 +08:00
ViperEkura 553a42702d refactor: replace diamond inheritance with mixin composition
- StreamStore/RecordStore → Streamable/Recordable (stateless mixins)
- Store is sole base class, no MRO ambiguity
- H5Store/MmapStore/JsonlStore mix in both traits explicitly
- segments_are_records declared per-subclass (H5/Jsonl=True, bin=False)
- Add tests for dpo_tokenize, lazy jsonl, dual-mode H5, stream-only bin
- Remove unused _to_tensor helper
2026-07-18 23:20:41 +08:00
ViperEkura b133fc9c07 refactor: split Store into StreamStore and RecordStore
- StreamStore: fetch(begin, end, key) for stream access (SEQ/SFT)
- RecordStore: mixin with fetch_record(i, key) for record access
- H5Store/MmapStore/JsonlStore now dual-inherit both (C3 MRO)
- JsonlStore supports lazy mode via processor= (no TokenizeTransform)
- RecordDataset base class holds processor, DPO/GRPO simplified
- dpo_tokenize pure function for on-the-fly JSONL tokenisation
- DatasetFactory builds processor for jsonl+record datasets
- train.py passes tokenizer_path=param_path uniformly
- progress: len(dataset) returns sample count (stream=windows, record=records)
- json no longer auto-detected as jsonl format
2026-07-18 23:04:31 +08:00
ViperEkura b33250dc28 refactor: decouple tokenizer from Store into Transform layer
- Extract tokenization/mask/position logic from JsonlStore into TokenizeTransform
- JsonlStore now pure reader: reads JSON records, delegates to transform
- Store no longer imports tokenizer or preprocessing components
- Replace per_record param with segments_are_records class attribute
- Store subclasses declare segment semantics as format-level property
2026-07-18 21:37:31 +08:00
ViperEkura a74e5b91a3 feat: add record-mode to Store for DPO/GRPO
- Store gains fetch_record/num_records alongside stream fetch/__len__
- save_bin/load_bin support per-record offsets via record_keys param
- H5Store/MmapStore/JsonlStore all support dual stream+record access
- DPODataset/GRPODataset use fetch_record, no cross-record concat
- dpo_collate_fn + collate_fn wired through TrainConfig
- fixes attention context leakage in DPO from windowed concatenation
2026-07-18 21:02:29 +08:00
ViperEkura 28886e4241 fix: make system prompt optional across scripts
- stream_chat: default empty system_prompt, single-turn mode
- generate_batch: drop hardcoded system role
- generate.py: preserve original fields in messages branch
  and use response_key for the output column name
2026-07-18 14:10:37 +08:00
ViperEkura 9d3ccfdffc fix: incremental decode to avoid U+FFFD in streaming
- StreamDecoder buffers incomplete multi-byte sequences
- Task.decode_new_token replaces per-token decode in scheduler
- flush_remaining emits final buffered text on task finish
2026-07-18 13:05:36 +08:00
ViperEkura a24a7b4da5 perf: merge decode batch for 10x throughput
- merge all active decode tasks into single forward pass (was grouped by next_pos)
- add per-task write_positions to ContiguousCacheView for correct KV writes
- override ContiguousCache.task_cached (base returned 0, caused prefill loops)
- add --cache_len/--frequency_penalty/--rep_window to generate.py
- chunked batch processing with tqdm progress

bench (1.2B model, 128 prompts, 64 tok, batch=128):
  before: 77.2s, ~111 tok/s
  after:   7.1s, ~1210 tok/s (10.9x)
2026-07-18 08:50:46 +08:00
ViperEkura f7df02f9a3 feat: add --num_samples to batch generation script 2026-07-18 01:13:38 +08:00
ViperEkura ee450686f3 fix: add option permutation to MMLU eval
- Few-shot examples now include subject preamble (consistent format)
- Add --seed flag for option permutation (default 0, -1 to disable)
- Shuffles A/B/C/D positions per-question to neutralise positional bias
2026-07-18 00:14:34 +08:00
ViperEkura 2565755e45 refactor: switch eval datasets to HuggingFace source
- Replace GitHub/berkeley direct downloads with HF datasets API
- MMLU: cais/mmlu (all config), map val->validation split, write per-subject CSV
- HumanEval: openai/openai_humaneval
- IFEval: google/IFEval
- Enables HF_ENDPOINT mirror for faster downloads in CN
2026-07-18 00:09:31 +08:00
97 changed files with 6359 additions and 2041 deletions
+25 -4
View File
@@ -23,6 +23,7 @@ jobs:
with: with:
name: pure-wheel name: pure-wheel
path: dist/*.whl path: dist/*.whl
if-no-files-found: error
build-cuda-linux: build-cuda-linux:
name: Build CUDA wheel (Linux) name: Build CUDA wheel (Linux)
@@ -50,6 +51,7 @@ jobs:
with: with:
name: cuda-wheel-linux name: cuda-wheel-linux
path: dist/*.whl path: dist/*.whl
if-no-files-found: error
release: release:
name: Attach wheels to release name: Attach wheels to release
@@ -58,14 +60,33 @@ jobs:
permissions: permissions:
contents: write contents: write
steps: steps:
- uses: actions/download-artifact@v4 - name: Download pure-Python wheel
uses: actions/download-artifact@v4
with: with:
pattern: "*-wheel" name: pure-wheel
merge-multiple: true path: release-assets/pure
- name: Download CUDA wheel
uses: actions/download-artifact@v4
with:
name: cuda-wheel-linux
path: release-assets/cuda
- name: Verify release assets
shell: bash
run: |
set -euo pipefail
pure_wheels=(release-assets/pure/*.whl)
cuda_wheels=(release-assets/cuda/*.whl)
test "${#pure_wheels[@]}" -eq 1
test "${#cuda_wheels[@]}" -eq 1
test "$(basename "${pure_wheels[0]}")" != "$(basename "${cuda_wheels[0]}")"
- name: Create release & upload assets - name: Create release & upload assets
uses: softprops/action-gh-release@v2 uses: softprops/action-gh-release@v2
with: with:
files: ./*.whl files: |
release-assets/pure/*.whl
release-assets/cuda/*.whl
tag_name: ${{ github.ref_name }} tag_name: ${{ github.ref_name }}
generate_release_notes: true generate_release_notes: true
+2 -2
View File
@@ -20,7 +20,7 @@
<a href="assets/docs/README-zh-CN.md">中文</a> • <a href="assets/docs/README-zh-CN.md">中文</a> •
<a href="https://github.com/ViperEkura/AstrAI/issues">Issue Tracker</a> • <a href="https://github.com/ViperEkura/AstrAI/issues">Issue Tracker</a> •
<a href="https://github.com/ViperEkura/AstrAI/discussions">Discussions</a> • <a href="https://github.com/ViperEkura/AstrAI/discussions">Discussions</a> •
<a href="https://huggingface.co/ViperEk/">HuggingFace</a> <a href="https://huggingface.co/ViperEkura">HuggingFace</a>
</div> </div>
<br> <br>
@@ -241,7 +241,7 @@ For major changes, please open an issue first to discuss what you would like to
- **GitHub Issues**: [Issue Tracker](https://github.com/ViperEkura/AstrAI/issues) - **GitHub Issues**: [Issue Tracker](https://github.com/ViperEkura/AstrAI/issues)
- **Discussions**: [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions) - **Discussions**: [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions)
- **HuggingFace**: [Model Hub](https://huggingface.co/ViperEk) - **HuggingFace**: [Model Hub](https://huggingface.co/ViperEkura)
### License ### License
+2 -2
View File
@@ -27,7 +27,7 @@
<a href="#chinese">中文</a> • <a href="#chinese">中文</a> •
<a href="https://github.com/ViperEkura/AstrAI/issues">问题追踪</a> • <a href="https://github.com/ViperEkura/AstrAI/issues">问题追踪</a> •
<a href="https://github.com/ViperEkura/AstrAI/discussions">讨论区</a> • <a href="https://github.com/ViperEkura/AstrAI/discussions">讨论区</a> •
<a href="https://huggingface.co/ViperEk">HuggingFace</a> <a href="https://huggingface.co/ViperEkura">HuggingFace</a>
</div> </div>
<br> <br>
@@ -247,7 +247,7 @@ SSE 流式格式、错误码和统计端点详见[推理文档](./inference.md)
- **GitHub Issues**: [问题追踪](https://github.com/ViperEkura/AstrAI/issues) - **GitHub Issues**: [问题追踪](https://github.com/ViperEkura/AstrAI/issues)
- **Discussions**: [GitHub 讨论区](https://github.com/ViperEkura/AstrAI/discussions) - **Discussions**: [GitHub 讨论区](https://github.com/ViperEkura/AstrAI/discussions)
- **HuggingFace**: [模型中心](https://huggingface.co/ViperEk) - **HuggingFace**: [模型中心](https://huggingface.co/ViperEkura)
### 许可证 ### 许可证
+288 -53
View File
@@ -28,17 +28,17 @@ classDiagram
class AutoRegressiveLMConfig { class AutoRegressiveLMConfig {
+Optional[int] vocab_size +Optional[int] vocab_size
+Optional[int] dim +Optional[int] hidden_size
+Optional[int] n_layers +Optional[int] num_hidden_layers
+Optional[float] norm_eps +Optional[float] rms_norm_eps
+Optional[int] dim_ffn +Optional[int] intermediate_size
+Optional[bool] tie_weight +Optional[bool] tie_word_embeddings
+Optional[dict] rope_scaling +Optional[dict] rope_scaling
+Optional[int] max_len +Optional[int] max_position_embeddings
+Optional[float] rope_theta +Optional[float] rope_theta
+str attn_type +str attn_type
+Optional[int] n_heads +Optional[int] num_attention_heads
+Optional[int] n_kv_heads +Optional[int] num_key_value_heads
+Optional[bool] use_qk_norm +Optional[bool] use_qk_norm
+Optional[bool] use_gated_attention +Optional[bool] use_gated_attention
+Optional[int] kv_lora_rank +Optional[int] kv_lora_rank
@@ -53,15 +53,15 @@ classDiagram
class EncoderConfig { class EncoderConfig {
+Optional[int] vocab_size +Optional[int] vocab_size
+Optional[int] dim +Optional[int] hidden_size
+Optional[int] n_layers +Optional[int] num_hidden_layers
+Optional[float] norm_eps +Optional[float] rms_norm_eps
+Optional[int] dim_ffn +Optional[int] intermediate_size
+Optional[int] max_len +Optional[int] max_position_embeddings
+Optional[float] rope_theta +Optional[float] rope_theta
+str attn_type +str attn_type
+Optional[int] n_heads +Optional[int] num_attention_heads
+Optional[int] n_kv_heads +Optional[int] num_key_value_heads
+Optional[bool] use_qk_norm +Optional[bool] use_qk_norm
+str ffn_type +str ffn_type
+Optional[dict] rope_scaling +Optional[dict] rope_scaling
@@ -117,7 +117,7 @@ classDiagram
+int n_epoch +int n_epoch
+int batch_per_device +int batch_per_device
+int grad_accum_steps +int grad_accum_steps
+float max_grad_norm +Optional[float] max_grad_norm
+list gradient_checkpointing_modules +list gradient_checkpointing_modules
+int start_epoch +int start_epoch
+int start_samples +int start_samples
@@ -141,6 +141,12 @@ classDiagram
+int val_step +int val_step
+float neftune_alpha +float neftune_alpha
+str parallel_mode +str parallel_mode
+int rollout_interval
+float rollout_temperature
+int rollout_top_k
+float rollout_top_p
+int rollout_max_tokens
+Optional[Callable] reward_model_fn
+dict executor_kwargs +dict executor_kwargs
+dict extra_kwargs +dict extra_kwargs
+validate() +validate()
@@ -177,13 +183,26 @@ classDiagram
class Store { class Store {
+Dict[str, List[Tensor]] _data +Dict[str, List[Tensor]] _data
+Dict[str, List[int]] _cum +Dict[str, List[int]] _cum
+Dict[str, List[int]] _offsets
+int _length +int _length
+int _num_records
+keys (property) +keys (property)
+load(path) +load(path)
+fetch(begin, end, keys)
+__len__() +__len__()
-_fetch_key(key, begin, end) Tensor -_normalize(raw, offsets)
-_normalize(raw) }
class Streamable {
<<mixin>>
+fetch(begin, end, keys)
-_fetch_stream_key(key, begin, end) Tensor
}
class Recordable {
<<mixin>>
+num_records (property)
+fetch_record(index, keys)
-_fetch_record_key(key, index) Tensor
} }
class H5Store { class H5Store {
@@ -195,7 +214,19 @@ classDiagram
+load(path) +load(path)
} }
class ResumableDistributedSampler { class JsonlStore {
+JsonlSource _source
+Callable _processor
+load(path, transform, processor)
+fetch_record(index, keys)
}
class JsonlSource {
+Path path
+load() List[dict]
}
class RDSampler {
+int epoch +int epoch
+int iter +int iter
} }
@@ -210,7 +241,7 @@ classDiagram
+Dict _entries +Dict _entries
+register(name) decorator +register(name) decorator
+create(train_type, window_size, stride) BaseDataset +create(train_type, window_size, stride) BaseDataset
+load(train_type, load_path, window_size, stride, storage_type) BaseDataset +load(train_type, load_path, window_size, stride, storage_type, tokenizer_path, max_len, store) BaseDataset
} }
} }
@@ -358,19 +389,103 @@ classDiagram
+forward(x) Tensor +forward(x) Tensor
+set_neftune_alpha(alpha) +set_neftune_alpha(alpha)
} }
class LoRAConfig {
+int r
+int alpha
+tuple target_modules
}
class LoRALinear {
+Linear weight
+Parameter lora_A, lora_B
+forward(x) Tensor
+merge()
}
} }
namespace preprocessing { namespace preprocessing {
class SectionRenderer {
+process_sections(item, sections, config, tokenizer) Tuple
+process_list_field(item, sections, config, tokenizer) Tuple
}
class BaseMaskBuilder { class BaseMaskBuilder {
<<abstract>> <<abstract>>
+build(item, config, tokenizer) Optional[dict] +build(item, config, tokenizer) Optional[dict]
} }
class SectionedMaskBuilder { class SingleOutputMaskBuilder {
+SectionRenderer renderer +SectionRenderer renderer
+build(item, config, tokenizer) Optional[dict] +build(item, config, tokenizer) Optional[dict]
+_build_single(item, config, tokenizer) Optional[dict] }
+_build_multi(item, sources_spec, config, tokenizer) Optional[dict]
class MultiOutputMaskBuilder {
+SectionRenderer renderer
+build(item, config, tokenizer) Optional[dict]
}
class SectionedMaskBuilder {
+build(item, config, tokenizer) Optional[dict]
}
class PackingStrategy {
<<abstract>>
+apply(keys, max_packed_len, truncation_mode) Dict
}
class PackingStrategyFactory {
+create(name, *args, **kwargs) PackingStrategy
}
class SimplePacking {
+apply(keys, max_packed_len, truncation_mode) Dict
}
class BFDPacking {
+apply(keys, max_packed_len, truncation_mode) Dict
}
class BFDSplitPacking {
+apply(keys, max_packed_len, truncation_mode) Dict
}
class PositionIdStrategy {
<<abstract>>
+generate(sequences) List[int]
}
class PositionIdStrategyFactory {
+create(name, *args, **kwargs) PositionIdStrategy
}
class NoPositionId {
+generate(sequences) List[int]
}
class DocResetPositionId {
+generate(sequences) List[int]
}
class ContinuousPositionId {
+generate(sequences) List[int]
}
class StoreWriter {
<<abstract>>
+save(output_dir, domain, shard_idx, tensors)
}
class StoreWriterFactory {
+create(name, *args, **kwargs) StoreWriter
}
class BinWriter {
+save(output_dir, domain, shard_idx, tensors)
}
class H5Writer {
+save(output_dir, domain, shard_idx, tensors)
} }
class Pipeline { class Pipeline {
@@ -378,6 +493,7 @@ classDiagram
+List[str] paths +List[str] paths
+str output_dir +str output_dir
+str tokenizer_path +str tokenizer_path
+AutoTokenizer tokenizer
+BaseMaskBuilder mask_builder +BaseMaskBuilder mask_builder
+PackingStrategy _packer +PackingStrategy _packer
+PositionIdStrategy _position_id +PositionIdStrategy _position_id
@@ -385,6 +501,18 @@ classDiagram
+transform(item) Optional[dict] +transform(item) Optional[dict]
+run() +run()
+_flush(domains, shard_idx) +_flush(domains, shard_idx)
+_inject_doc_reset_position_ids(keys, mode, seqs) Dict
+_inject_continuous_position_ids(tensors, mode, seqs) Dict
+_to_tensors(keys) Dict
}
class TokenizeTransform {
+PipelineConfig config
+AutoTokenizer tokenizer
+BaseMaskBuilder mask_builder
+PositionIdStrategy position_strategy
+from_config_file(path) TokenizeTransform
+apply(records) Dict[str, list]
} }
} }
@@ -457,7 +585,7 @@ classDiagram
class TrainContextBuilder { class TrainContextBuilder {
+TrainConfig config +TrainConfig config
+with_resume_dir(resume_dir) TrainContextBuilder +with_param_path(param_path, resume) TrainContextBuilder
+build() TrainContext +build() TrainContext
} }
@@ -495,13 +623,39 @@ classDiagram
} }
class GRPOStrategy { class GRPOStrategy {
+nn.Module old_model
+nn.Module ref_model +nn.Module ref_model
+float clip_eps +float clip_eps
+float kl_coef +float kl_coef
+int group_size +int group_size
+int sync_interval
+compute_loss(batch) Tensor +compute_loss(batch) Tensor
+sync_ref_model() +sync_old_model()
}
class RawRollout {
+Tensor prompts
+Tensor responses
+Tensor response_mask
+Tensor logprobs_old
}
class RolloutResult {
+Tensor rewards
}
class BaseRewardModel {
<<abstract>>
+score(prompts, responses) Tensor
}
class RolloutGenerator {
+generate(batch) RawRollout
}
class RolloutRunner {
+step()
+clear_cache()
+__call__(batch) Tuple[RolloutResult, bool]
} }
class BaseScheduler { class BaseScheduler {
@@ -551,7 +705,7 @@ classDiagram
} }
class GradientClippingCallback { class GradientClippingCallback {
+float max_grad_norm +Optional[float] max_grad_norm
+on_optimizer_step(context) +on_optimizer_step(context)
} }
@@ -817,12 +971,21 @@ classDiagram
+apply(logits, filter_value) Tensor +apply(logits, filter_value) Tensor
} }
class FrequencyPenaltyStrategy {
+float penalty
+apply(logits, filter_value, input_ids, input_mask) Tensor
}
class SamplingPipeline { class SamplingPipeline {
+List[BaseSamplingStrategy] strategies +List[BaseSamplingStrategy] strategies
+apply(logits, filter_value) Tensor +apply(logits, filter_value) Tensor
+sample(logits, filter_value) Tensor +sample(logits, filter_value) Tensor
} }
class StreamDecoder {
+push(token_id) str
}
class GenerateResult { class GenerateResult {
+List[Tuple[int, str]] tokens +List[Tuple[int, str]] tokens
+List[str] results +List[str] results
@@ -841,6 +1004,17 @@ classDiagram
+Optional[str] tool_call_id +Optional[str] tool_call_id
} }
class FunctionDef {
+str name
+Optional[str] description
+Optional[Dict] parameters
}
class ToolDef {
+str type
+FunctionDef function
}
class ChatCompletionRequest { class ChatCompletionRequest {
+str model +str model
+List[ChatMessage] messages +List[ChatMessage] messages
@@ -929,9 +1103,20 @@ classDiagram
+str yielded +str yielded
} }
class get_app { class BaseToolParser {
<<module>> <<abstract>>
+get_app() FastAPI +feed(body, current_token_ids, delta_token_ids) List[Dict]
+parse_complete(body) Optional[Dict]
+has_tool_calls (property) bool
}
class ToolParserFactory {
+create(name, *args, **kwargs) BaseToolParser
}
class SimpleJsonToolParser {
+feed(body, current_token_ids, delta_token_ids) List[Dict]
+parse_complete(body) Optional[Dict]
} }
} }
@@ -954,14 +1139,17 @@ classDiagram
} }
namespace parallel { namespace parallel {
class setup { class LaunchStrategy {
<<module>> <<abstract>>
+spawn_parallel_fn(func, world_size, backend, master_addr, master_port, device_type, start_method, **kwargs) +launch(func, **kwargs)
+setup_parallel(rank, world_size, backend, master_addr, master_port, device_type) contextmanager }
+get_current_device() str
+get_world_size() int class TorchrunStrategy {
+get_rank() int +launch(func, **kwargs)
+only_on_rank(rank, sync=False) decorator }
class LocalStrategy {
+launch(func, **kwargs)
} }
class GradientState { class GradientState {
@@ -990,7 +1178,7 @@ classDiagram
class BaseExecutor { class BaseExecutor {
+GradientState gradient_state +GradientState gradient_state
+prepare(model, optimizer, dataloader, scheduler) tuple +prepare(model_fn, optimizer_fn, scheduler_fn, before_wrap) tuple
+accumulate(model) context manager +accumulate(model) context manager
+backward(loss) +backward(loss)
+unwrap_model(model) dict +unwrap_model(model) dict
@@ -1012,6 +1200,12 @@ classDiagram
+unwrap_model(model) dict +unwrap_model(model) dict
} }
class FSDP2Executor {
-_prepare_model(model) nn.Module
-_no_sync(model) context manager
+unwrap_model(model) dict
}
class ExecutorFactory { class ExecutorFactory {
+Dict _entries +Dict _entries
+register(name) decorator +register(name) decorator
@@ -1069,9 +1263,16 @@ classDiagram
Store <|-- H5Store Store <|-- H5Store
Store <|-- MmapStore Store <|-- MmapStore
Store <|-- JsonlStore Store <|-- JsonlStore
H5Store --|> Streamable
H5Store --|> Recordable
MmapStore --|> Streamable
MmapStore --|> Recordable
JsonlStore --|> Streamable
JsonlStore --|> Recordable
BaseSamplingStrategy <|-- TemperatureStrategy BaseSamplingStrategy <|-- TemperatureStrategy
BaseSamplingStrategy <|-- TopKStrategy BaseSamplingStrategy <|-- TopKStrategy
BaseSamplingStrategy <|-- TopPStrategy BaseSamplingStrategy <|-- TopPStrategy
BaseSamplingStrategy <|-- FrequencyPenaltyStrategy
ParallelModel <|-- RowParallelLinear ParallelModel <|-- RowParallelLinear
ParallelModel <|-- ColumnParallelLinear ParallelModel <|-- ColumnParallelLinear
AutoModel <|-- AutoRegressiveLM AutoModel <|-- AutoRegressiveLM
@@ -1095,12 +1296,31 @@ classDiagram
BaseFactory <|-- ExecutorFactory BaseFactory <|-- ExecutorFactory
BaseFactory <|-- ConfigFactory BaseFactory <|-- ConfigFactory
BaseFactory <|-- MaskBuilderFactory BaseFactory <|-- MaskBuilderFactory
BaseFactory <|-- PackingStrategyFactory
BaseFactory <|-- PositionIdStrategyFactory
BaseFactory <|-- StoreWriterFactory
BaseFactory <|-- ToolParserFactory
BaseExecutor <|-- NoneExecutor BaseExecutor <|-- NoneExecutor
BaseExecutor <|-- DDPExecutor BaseExecutor <|-- DDPExecutor
BaseExecutor <|-- FSDPExecutor BaseExecutor <|-- FSDPExecutor
BaseExecutor <|-- FSDP2Executor
ResponseBuilder <|-- OpenAIResponseBuilder ResponseBuilder <|-- OpenAIResponseBuilder
ResponseBuilder <|-- AnthropicResponseBuilder ResponseBuilder <|-- AnthropicResponseBuilder
BaseToolParser <|-- SimpleJsonToolParser
BaseMaskBuilder <|-- SectionedMaskBuilder BaseMaskBuilder <|-- SectionedMaskBuilder
BaseMaskBuilder <|-- SingleOutputMaskBuilder
BaseMaskBuilder <|-- MultiOutputMaskBuilder
PackingStrategy <|-- SimplePacking
PackingStrategy <|-- BFDPacking
BFDPacking <|-- BFDSplitPacking
PositionIdStrategy <|-- NoPositionId
PositionIdStrategy <|-- DocResetPositionId
PositionIdStrategy <|-- ContinuousPositionId
StoreWriter <|-- BinWriter
StoreWriter <|-- H5Writer
RawRollout <|-- RolloutResult
LaunchStrategy <|-- TorchrunStrategy
LaunchStrategy <|-- LocalStrategy
KVCache <|-- PageCache KVCache <|-- PageCache
KVCache <|-- ContiguousCache KVCache <|-- ContiguousCache
CacheView <|-- PageCacheView CacheView <|-- PageCacheView
@@ -1122,6 +1342,8 @@ classDiagram
EmbeddingEncoder *-- Embedding EmbeddingEncoder *-- Embedding
DecoderBlock *-- RMSNorm DecoderBlock *-- RMSNorm
ChatCompletionRequest *-- ChatMessage ChatCompletionRequest *-- ChatMessage
ChatCompletionRequest *-- ToolDef
ToolDef *-- FunctionDef
MessagesRequest *-- AnthropicMessage MessagesRequest *-- AnthropicMessage
BaseExecutor *-- GradientState BaseExecutor *-- GradientState
AccumOptimizer o-- GradientState AccumOptimizer o-- GradientState
@@ -1143,11 +1365,20 @@ classDiagram
BaseDataset o-- Store BaseDataset o-- Store
Pipeline o-- PipelineConfig Pipeline o-- PipelineConfig
Pipeline o-- BaseMaskBuilder Pipeline o-- BaseMaskBuilder
Pipeline o-- AutoTokenizer
Pipeline o-- PackingStrategy
Pipeline o-- PositionIdStrategy
Pipeline o-- StoreWriter
TokenizeTransform o-- AutoTokenizer
TokenizeTransform o-- BaseMaskBuilder
%% --- Dependency (uses temporarily) --- %% --- Dependency (uses temporarily) ---
TrainConfig ..> BaseStrategy : selects TrainConfig ..> BaseStrategy : selects
PipelineConfig ..> MaskBuilderFactory : selects PipelineConfig ..> MaskBuilderFactory : selects
MaskBuilderFactory ..> BaseMaskBuilder : creates MaskBuilderFactory ..> BaseMaskBuilder : creates
PackingStrategyFactory ..> PackingStrategy : creates
PositionIdStrategyFactory ..> PositionIdStrategy : creates
StoreWriterFactory ..> StoreWriter : creates
StrategyFactory ..> BaseStrategy : creates StrategyFactory ..> BaseStrategy : creates
SchedulerFactory ..> BaseScheduler : creates SchedulerFactory ..> BaseScheduler : creates
DatasetFactory ..> BaseDataset : creates DatasetFactory ..> BaseDataset : creates
@@ -1166,12 +1397,13 @@ classDiagram
ExecutorFactory ..> NoneExecutor : creates ExecutorFactory ..> NoneExecutor : creates
ExecutorFactory ..> DDPExecutor : creates ExecutorFactory ..> DDPExecutor : creates
ExecutorFactory ..> FSDPExecutor : creates ExecutorFactory ..> FSDPExecutor : creates
ExecutorFactory ..> FSDP2Executor : creates
ToolParserFactory ..> BaseToolParser : creates
TrainContextBuilder ..> ExecutorFactory : creates TrainContextBuilder ..> ExecutorFactory : creates
Trainer ..> TrainContextBuilder : uses Trainer ..> TrainContextBuilder : uses
TrainContextBuilder ..> TrainContext : creates TrainContextBuilder ..> TrainContext : creates
Trainer ..> Functions : spawns
TrainContextBuilder ..> StrategyFactory : uses TrainContextBuilder ..> StrategyFactory : uses
TrainContextBuilder ..> ResumableDistributedSampler : creates TrainContextBuilder ..> RDSampler : creates
Checkpoint ..> Checkpoint : serializes Checkpoint ..> Checkpoint : serializes
CheckpointCallback ..> Checkpoint : creates CheckpointCallback ..> Checkpoint : creates
PageCache ..> PageCacheView : binds PageCache ..> PageCacheView : binds
@@ -1182,11 +1414,14 @@ classDiagram
AnthropicResponseBuilder ..> MessagesRequest : receives AnthropicResponseBuilder ..> MessagesRequest : receives
ProtocolHandler ..> StopChecker : creates ProtocolHandler ..> StopChecker : creates
ProtocolHandler ..> GenContext : creates ProtocolHandler ..> GenContext : creates
RolloutGenerator ..> InferenceScheduler : uses
RolloutRunner ..> RolloutGenerator : uses
RolloutRunner ..> BaseRewardModel : uses
%% --- Association (general usage) --- %% --- Association (general usage) ---
Trainer --> TrainConfig Trainer --> TrainConfig
DPOStrategy --> AutoModel DPOStrategy --> AutoModel
GRPOStrategy --> AutoModel GRPOStrategy --> AutoModel : policy/old/ref
InferenceScheduler --> Task InferenceScheduler --> Task
InferenceScheduler --> TaskStatus InferenceScheduler --> TaskStatus
Task --> TaskStatus Task --> TaskStatus
@@ -1203,14 +1438,14 @@ classDiagram
| Module | Components | Description | | 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.config** | BaseConfig, BaseModelConfig, AutoRegressiveLMConfig, EncoderConfig, ConfigFactory, TrainConfig, PipelineConfig, InputConfig, ProcessingConfig, OutputConfig | Configuration management (to_dict/from_dict, to_file/from_file) |
| **astrai.preprocessing** | BaseMaskBuilder, MaskBuilderFactory, SectionedMaskBuilder, Pipeline, filter_by_length, PackingStrategy, PackingStrategyFactory, PositionIdStrategy, PositionIdStrategyFactory, StoreWriter, StoreWriterFactory | Declarative JSON-driven data preprocessing | | **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** | BaseDatasetGRPODataset, StoreJsonlStore/MmapStore/H5Store, StoreFactory, ResumableDistributedSampler, DatasetFactory | Dataset loading and management | | **astrai.dataset** | BaseDataset, SEQDataset, SFTDataset, DPODataset, GRPODataset, Store, Streamable, Recordable, H5Store, MmapStore, JsonlSource, JsonlStore, StoreFactory, RDSampler, DatasetFactory | Dataset loading and management |
| **astrai.serialization** | Checkpoint | Model serialization | | **astrai.serialization** | Checkpoint | Model serialization |
| **astrai.model** | AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, RotaryEmbedding, Embedding | Neural network model | | **astrai.model** | AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, LoRAConfig, LoRALinear, RotaryEmbedding, Embedding | Neural network model |
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template | | **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategyGRPOStrategy, StrategyFactory, BaseSchedulerWSDScheduler, SchedulerFactory, TrainCallback(Protocol)MetricCallback, CallbackFactory | Training workflow | | **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategyGRPOStrategy, StrategyFactory, BaseSchedulerWSDScheduler, SchedulerFactory, TrainCallback(Protocol)MetricCallback, CallbackFactory, RawRollout, RolloutResult, BaseRewardModel, RolloutGenerator, RolloutRunner | Training workflow |
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, KVCacheContiguousCache/PageCache, CacheViewContiguousCacheView/PageCacheView, AllocatorStorage, Task, TaskManager, TaskStatus, GenerationRequest, GenerateResult, BaseSamplingStrategySamplingPipeline, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, ChatMessageMessagesRequest, app | Inference service | | **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, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, FSDPExecutor, GradientState, AccumOptimizer, AccumScheduler, ParallelModel, RowParallelLinear, ColumnParallelLinear | Distributed parallel & gradient accumulation | | **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, LaunchStrategy, TorchrunStrategy, LocalStrategy, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, FSDPExecutor, FSDP2Executor, GradientState, AccumOptimizer, AccumScheduler, ParallelModel, RowParallelLinear, ColumnParallelLinear | Distributed parallel & gradient accumulation |
| **astrai.factory** | BaseFactory | Component registration | | **astrai.factory** | BaseFactory | Component registration |
| **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers | | **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers |
@@ -1218,7 +1453,7 @@ classDiagram
| Pattern | Classes | Purpose | | Pattern | Classes | Purpose |
|---------|---------|---------| |---------|---------|---------|
| **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory`, `MaskBuilderFactory`, `StoreWriterFactory`, `PackingStrategyFactory`, `PositionIdStrategyFactory` | Decorator-based component creation | | **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory`, `MaskBuilderFactory`, `StoreWriterFactory`, `PackingStrategyFactory`, `PositionIdStrategyFactory`, `ToolParserFactory` | Decorator-based component creation |
| **Registry** | `BaseFactory` | Component registration | | **Registry** | `BaseFactory` | Component registration |
| **Strategy** | `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy` | Training strategy switching | | **Strategy** | `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy` | Training strategy switching |
| **Strategy (Sampling)** | `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations | | **Strategy (Sampling)** | `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations |
@@ -1227,7 +1462,7 @@ classDiagram
| **Observer** | `TrainCallback`, callback implementations | Training process monitoring | | **Observer** | `TrainCallback`, callback implementations | Training process monitoring |
| **Context** | `TrainContext` | Unified training state bag | | **Context** | `TrainContext` | Unified training state bag |
| **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction | | **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction |
| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor` | Gradient accumulation & model distribution | | **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor`, `FSDP2Executor` | Gradient accumulation & model distribution |
| **Storage** | `Store`, `H5Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support | | **Storage** | `Store`, `H5Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support |
| **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching | | **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching |
| **AutoModel Registry** | `AutoModel`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading | | **AutoModel Registry** | `AutoModel`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading |
@@ -1237,7 +1472,7 @@ classDiagram
1. **Config → Training**: `TrainConfig` holds `model_fn`, `dataset`, `optimizer_fn`, `scheduler_fn`, `parallel_mode`, `executor_kwargs` 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 2. **Training Flow**: `Trainer``TrainContextBuilder``TrainContext`, uses `BaseStrategy` for loss, `BaseExecutor` for gradient accumulation + model distribution
3. **Strategy Selection**: `StrategyFactory` creates strategy by `train_type` 3. **Strategy Selection**: `StrategyFactory` creates strategy by `train_type`
4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)``NoneExecutor` / `DDPExecutor` / `FSDPExecutor` 4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)``NoneExecutor` / `DDPExecutor` / `FSDPExecutor` / `FSDP2Executor`
5. **Inference Flow**: `InferenceEngine``InferenceScheduler``AutoRegressiveLM`, backed by `KVCache` + `SamplingPipeline` 5. **Inference Flow**: `InferenceEngine``InferenceScheduler``AutoRegressiveLM`, backed by `KVCache` + `SamplingPipeline`
6. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP 6. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP
7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (H5Store/MmapStore/JsonlStore) loads data with explicit `_length` and multi-segment `_data` 7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (H5Store/MmapStore/JsonlStore) loads data with explicit `_length` and multi-segment `_data`
@@ -1246,4 +1481,4 @@ classDiagram
10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops 10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops
11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers 11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers
> Document Update Time: 2026-07-09 > Document Update Time: 2026-07-20
+34 -16
View File
@@ -61,41 +61,59 @@ StoreFactory.create("bin") → MmapStore
StoreFactory.create("jsonl") → JsonlStore StoreFactory.create("jsonl") → JsonlStore
``` ```
**H5Store**: Reads HDF5 files. Tensors are loaded into host memory and normalized into segmented storage. 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.
**MmapStore**: Memory-maps `.bin` files. OS page cache sharing is native — no explicit `share_memory_()` needed. Uses `torch.from_numpy(np.memmap(...))`. **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.
**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. **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).
All backends normalise tensors into `Store._data[Dict[str, List[Tensor]]]` + `Store._cum[Dict[str, List[int]]]` (cumulative lengths for bisect-based indexing). **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 ## Data Keys by Training Type
| Type | Storage Keys | | Type | Storage Keys | Access Mode |
|------|-------------| |------|-------------|-------------|
| `seq` | `sequence` (→ input_ids, target_ids via offset-by-1) | | `seq` | `sequence` (→ input_ids, target_ids via offset-by-1) | stream (`fetch`) |
| `sft` | `sequence`, `loss_mask`, `position_ids` | | `sft` | `sequence`, `loss_mask`, `position_ids` | stream (`fetch`) |
| `dpo` | `chosen`, `rejected`, `chosen_mask`, `rejected_mask` | | `dpo` | `chosen`, `rejected`, `chosen_mask`, `rejected_mask` | record (`fetch_record`) |
| `grpo` | `prompts`, `responses`, `masks`, `rewards` | | `grpo` | `prompts`, `responses`, `masks`, `rewards` | record (`fetch_record`) |
## Dataset Architecture ## Dataset Architecture
``` ```
DatasetFactory.load(train_type, load_path, window_size, stride=None, storage_type=None) 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) → BaseDataset.load(load_path, storage_type=None)
→ detect_format(load_path) → detect_format(load_path)
→ StoreFactory.create(storage_type) → StoreFactory.create(storage_type)
→ Store.load(load_path) → Store.load(load_path)
→ _normalize(raw) # base Store, shared by both backends → _normalize(raw) # base Store, shared by both backends
→ Store._data[Dict[str, List[Tensor]]] + _cum[Dict[str, List[int]]] → Store._data[Dict[str, List[Tensor]]]
→ BaseDataset.__getitem__(idx) + _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) → get_index(idx) → [begin, end)
→ Store.fetch(begin, end, keys) → Tensor / Dict[str, Tensor] → 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]
``` ```
`window_size` = max input length, `stride` = step between consecutive samples (defaults to `window_size`, optional). `storage_type` defaults to `None` (auto-detect via `detect_format`). Class hierarchy: `BaseDataset``SEQDataset` / `SFTDataset` (stream); `BaseDataset``RecordDataset``DPODataset` / `GRPODataset` (record).
`Store.fetch(begin, end, keys)` 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()`. `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 ## Sampler
@@ -109,4 +127,4 @@ DatasetFactory.load(train_type, load_path, window_size, stride=None, storage_typ
Standard PyTorch `DataLoader` with configurable `batch_size`, `num_workers`, `pin_memory`, `prefetch_factor`. Sampler produces indices; dataloader fetches tensor batches via `__getitem__`. 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-09 > Document Update Time: 2026-07-19
+2 -2
View File
@@ -32,7 +32,7 @@ ContiguousCache (simple contiguous per-slot cache)
├── ContiguousCacheView bundles k/v tensors + slot indices for attention layers ├── ContiguousCacheView bundles k/v tensors + slot indices for attention layers
``` ```
Created by default when no cache is passed to `InferenceScheduler`. Each task occupies a fixed slot of `[max_seq_len, n_kv_heads, head_dim]`. Simple and efficient for small-to-medium batch sizes. 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.
### PageCache (paged with prefix sharing) ### PageCache (paged with prefix sharing)
@@ -42,7 +42,7 @@ PageCache (paged KV cache with prefix sharing, alternative)
│ ├── Allocator bitmask-based page allocator + ref-count + LRU │ ├── Allocator bitmask-based page allocator + ref-count + LRU
│ └── PrefixCache hash-based prefix matching (page_hash via polynomial hash) │ └── PrefixCache hash-based prefix matching (page_hash via polynomial hash)
├── TaskTable maps task_id → page_table + cached token count ├── TaskTable maps task_id → page_table + cached token count
├── Storage k_cache / v_cache tensors (n_layers × n_pages × page_size × n_kv_heads × head_dim) ├── 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 └── PageCacheView bundles Storage + page_table + total_len for attention layers
``` ```
+21 -8
View File
@@ -13,7 +13,7 @@
| Parameter | Description | Default | | Parameter | Description | Default |
|-----------|-------------|---------| |-----------|-------------|---------|
| `--train_type` | Training type (`seq`, `sft`, `dpo`, `grpo`) | required | | `--train_type` | Training type (`seq`, `sft`, `dpo`, `grpo`, `online_grpo`, `online_dpo`) | required |
| `--data_root_path` | Dataset root directory | required | | `--data_root_path` | Dataset root directory | required |
| `--param_path` | Model parameters or checkpoint path | required | | `--param_path` | Model parameters or checkpoint path | required |
| `--n_epoch` | Total training epochs | 1 | | `--n_epoch` | Total training epochs | 1 |
@@ -26,7 +26,7 @@
|-----------|-------------|---------| |-----------|-------------|---------|
| `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 | | `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 |
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 | | `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
| `--max_grad_norm` | Maximum gradient norm for clipping | 1.0 | | `--max_grad_norm` | Maximum gradient norm for clipping (None disables) | 1.0 |
### Optimizer (MuonMix) ### Optimizer (MuonMix)
@@ -44,7 +44,7 @@ Combined optimizer: matrix parameters via **Muon**, non-matrix via **AdamW** (`f
| Parameter | Description | Default | | Parameter | Description | Default |
|-----------|-------------|---------| |-----------|-------------|---------|
| `--window_size` | Max input sequence length | model config `max_len` | | `--window_size` | Max input sequence length | model config `max_position_embeddings` |
| `--stride` | Stride for sliding window over sequences | None | | `--stride` | Stride for sliding window over sequences | None |
| `--random_seed` | Random seed for reproducibility | 3407 | | `--random_seed` | Random seed for reproducibility | 3407 |
| `--num_workers` | DataLoader worker processes | 4 | | `--num_workers` | DataLoader worker processes | 4 |
@@ -100,18 +100,31 @@ Combined optimizer: matrix parameters via **Muon**, non-matrix via **AdamW** (`f
| `--group_size` | GRPO group size | 4 | `grpo` | | `--group_size` | GRPO group size | 4 | `grpo` |
| `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo` | | `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo` |
| `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 | `grpo` | | `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 | `grpo` |
| `--grpo_sync_interval` | GRPO ref_model sync interval (steps) | 200 | `grpo` |
| `--neftune_alpha` | NEFTune noise alpha (0=disabled, typical: 5.0) | 0.0 | `sft` | | `--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
a `BaseRewardModel` factory in `TrainConfig`; `train.py` does not currently
provide a command-line option for configuring one.
| Parameter | Description | Default |
|-----------|-------------|---------|
| `--rollout_interval` | Optimizer steps between rollout refreshes | 512 |
| `--rollout_temperature` | Rollout sampling temperature | 0.7 |
| `--rollout_top_k` | Rollout top-k filtering (`0` disables) | 0 |
| `--rollout_top_p` | Rollout nucleus sampling threshold | 0.9 |
| `--rollout_max_tokens` | Maximum generated tokens per response | 1024 |
### Scheduler ### Scheduler
| Parameter | Description | Default | | Parameter | Description | Default |
|-----------|-------------|---------| |-----------|-------------|---------|
| `--schedule_type` | LR scheduler type (`cosine`, `sgdr`, `wsd`) | cosine | | `--schedule_type` | LR scheduler type (`cosine`, `sgdr`, `wsd`) | cosine |
| `--min_rate` | Minimum LR as fraction of base LR | None (scheduler default: 0.01) | | `--min_rate` | Minimum LR as fraction of base LR | None (scheduler default: 0.05 for cosine/SGDR, 0.0 for WSD) |
| `--cycle_length` | SGDR first cycle length in steps | None (total_steps - warmup_steps) | | `--cycle_length` | SGDR first cycle length in steps | None (total_steps - warmup_steps) |
| `--t_mult` | SGDR cycle length multiplier per restart | 2 | | `--t_mult` | SGDR cycle length multiplier per restart | 2 |
| `--stable_steps` | WSD stable plateau steps | None (required for wsd) | | `--stable_steps` | WSD stable plateau steps | None (80% of post-warmup steps) |
| `--decay_steps` | WSD decay steps | None (total_steps - warmup_steps - stable_steps) | | `--decay_steps` | WSD decay steps | None (total_steps - warmup_steps - stable_steps) |
### Usage Example ### Usage Example
@@ -173,7 +186,7 @@ See [Inference Guide](inference.md) for HTTP API documentation.
| `--top_k` | int | `30` | Top-k filtering | | `--top_k` | int | `30` | Top-k filtering |
| `--top_p` | float | `0.95` | Nucleus sampling threshold | | `--top_p` | float | `0.95` | Nucleus sampling threshold |
| `--batch_size` | int | `1` | Batch size for generation | | `--batch_size` | int | `1` | Batch size for generation |
| `--max_tokens` | int | model config `max_len` | Maximum tokens to generate | | `--max_tokens` | int | model config `max_position_embeddings` | Maximum tokens to generate |
Usage: Usage:
```bash ```bash
@@ -201,4 +214,4 @@ See [Preprocessing Guide](preprocessing.md) for config file format and examples.
--- ---
> Document Update Time: 2026-07-09 > Document Update Time: 2026-07-20
+29 -15
View File
@@ -6,7 +6,7 @@
- [Causal Mask](#causal-mask) - [Causal Mask](#causal-mask)
- [Rotary Position Embedding (RoPE)](#rotary-position-embedding-rope) - [Rotary Position Embedding (RoPE)](#rotary-position-embedding-rope)
- [Training Loop](#training-loop) - [Training Loop](#training-loop)
- [Strategies](#strategies) — SEQ, SFT, DPO, GRPO - [Strategies](#strategies) — SEQ, SFT, DPO, GRPO, online rollout
- [LR Schedulers](#lr-schedulers) - [LR Schedulers](#lr-schedulers)
- [Gradient Checkpointing](#gradient-checkpointing) - [Gradient Checkpointing](#gradient-checkpointing)
- [Checkpoint](#checkpoint) - [Checkpoint](#checkpoint)
@@ -86,7 +86,7 @@ on_train_end
| `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` | | `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` |
| `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricCallback`, `GradientCheckpointingCallback` | | `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricCallback`, `GradientCheckpointingCallback` |
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm), `gradient_clipping`. Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm), `gradient_clipping` (always registered; computes grad norm, clips only when `max_grad_norm` is not `None`).
## Strategies ## Strategies
@@ -118,7 +118,7 @@ $$
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] 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="mean"`. Keys: `chosen`, `rejected`, `chosen_mask`, `rejected_mask`. Parameters: `beta=0.1`, `reduction="sum"`. Keys: `chosen`, `rejected`, `chosen_mask`, `rejected_mask`.
### GRPO (Group Relative Policy Optimization) ### GRPO (Group Relative Policy Optimization)
@@ -135,13 +135,30 @@ $$
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] 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{ref}}(a_t|s_t)$ is the where $\rho_t = \pi_\theta(a_t|s_t) / \pi_{\text{old}}(a_t|s_t)$ is the
per-token probability ratio and the expectations are over valid response tokens. per-token importance sampling ratio against the behaviour policy
(`old_model`, synced externally between data-generation rounds) and the
expectations are over valid response tokens. The KL term regularises
$\pi_\theta$ towards a frozen reference model (`ref_model`, typically
the SFT checkpoint).
Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`, `sync_interval=200`. Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`. External sync of `old_model` weights via `sync_old_model()` between data-generation rounds.
Keys: `prompts`, `responses`, `masks`, `rewards`. Keys: `prompts`, `responses`, `masks`, `rewards`.
### Online Rollout
`online_grpo` and `online_dpo` use the respective GRPO and DPO strategies with
a `RolloutRunner`. The runner renders prompts through the tokenizer chat
template, generates grouped responses through `InferenceScheduler`, then scores
them with a `BaseRewardModel`. It refreshes cached rollouts every
`rollout_interval` optimizer steps. `online_grpo` synchronizes `old_model` when
a fresh rollout is produced.
Online strategies require `TrainConfig.reward_model_fn`. `train.py` exposes the
rollout sampling parameters but does not yet offer a CLI argument for the reward
model factory.
## LR Schedulers ## LR Schedulers
| Type | Class | Description | | Type | Class | Description |
@@ -158,6 +175,7 @@ Trades compute for memory by recomputing activations during backward pass. Speci
```python ```python
from astrai.model.components.decoder_block import DecoderBlock from astrai.model.components.decoder_block import DecoderBlock
config = TrainConfig(..., gradient_checkpointing_modules=[DecoderBlock]) config = TrainConfig(..., gradient_checkpointing_modules=[DecoderBlock])
``` ```
@@ -177,18 +195,14 @@ Model config (`context.model_config`) saved into `config.json` during training v
## TrainContextBuilder (Builder Pattern) ## TrainContextBuilder (Builder Pattern)
```python ```python
context = ( context = TrainContextBuilder(config).with_param_path(param_path, resume=True).build()
TrainContextBuilder(config)
.with_resume_dir(resume_dir)
.build()
)
# Returns TrainContext with model, strategy, optimizer, scheduler, dataloader, checkpoint # Returns TrainContext with model, strategy, optimizer, scheduler, dataloader, checkpoint
``` ```
- Loads checkpoint weights if provided - Loads checkpoint weights before the model is wrapped
- Creates executor via `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)` - Creates executor via `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)`
- Calls `executor.prepare(model, optimizer, dataloader, scheduler)` for model distribution (e.g. DDP) + gradient accumulation wrappers - Calls `executor.prepare(model_fn, optimizer_fn, scheduler_fn, before_wrap=...)`; the executor creates, wraps, then builds the optimizer and scheduler for the wrapped model
- Creates `ResumableDistributedSampler` for shuffle+resume - Creates `RDSampler` for shuffle+resume
- Builds strategy via `StrategyFactory.create(train_type, model, device, **kwargs)` - Builds strategy via `StrategyFactory.create(train_type, model, device, **kwargs)`
## Training CLI ## Training CLI
@@ -218,4 +232,4 @@ nohup python scripts/tools/train.py \
Full parameter reference at [params.md](params.md). Full parameter reference at [params.md](params.md).
> Document Update Time: 2026-07-09 > Document Update Time: 2026-07-20
+3 -3
View File
@@ -1,4 +1,4 @@
__version__ = "1.3.9" __version__ = "1.3.11"
__author__ = "ViperEkura" __author__ = "ViperEkura"
from astrai.config import ( from astrai.config import (
@@ -12,7 +12,7 @@ from astrai.config import (
from astrai.dataset import ( from astrai.dataset import (
BaseDataset, BaseDataset,
DatasetFactory, DatasetFactory,
ResumableDistributedSampler, RDSampler,
Store, Store,
StoreFactory, StoreFactory,
) )
@@ -77,7 +77,7 @@ __all__ = [
"Pipeline", "Pipeline",
"PipelineConfig", "PipelineConfig",
"ProtocolHandler", "ProtocolHandler",
"ResumableDistributedSampler", "RDSampler",
"SamplingPipeline", "SamplingPipeline",
"SchedulerFactory", "SchedulerFactory",
"Store", "Store",
+15 -15
View File
@@ -29,19 +29,19 @@ class AutoRegressiveLMConfig(BaseModelConfig):
"""Configuration for autoregressive language model.""" """Configuration for autoregressive language model."""
vocab_size: Optional[int] = None vocab_size: Optional[int] = None
dim: Optional[int] = None hidden_size: Optional[int] = None
n_layers: Optional[int] = None num_hidden_layers: Optional[int] = None
norm_eps: Optional[float] = None rms_norm_eps: Optional[float] = None
dim_ffn: Optional[int] = None intermediate_size: Optional[int] = None
tie_weight: Optional[bool] = None tie_word_embeddings: Optional[bool] = None
max_len: Optional[int] = None max_position_embeddings: Optional[int] = None
rope_theta: Optional[float] = None rope_theta: Optional[float] = None
rope_scaling: Optional[dict] = None rope_scaling: Optional[dict] = None
attn_type: str = "gqa" attn_type: str = "gqa"
n_heads: Optional[int] = None num_attention_heads: Optional[int] = None
n_kv_heads: Optional[int] = None num_key_value_heads: Optional[int] = None
use_qk_norm: Optional[bool] = None use_qk_norm: Optional[bool] = None
use_gated_attention: Optional[bool] = None use_gated_attention: Optional[bool] = None
@@ -62,18 +62,18 @@ class EncoderConfig(BaseModelConfig):
"""Configuration for embedding encoder model.""" """Configuration for embedding encoder model."""
vocab_size: Optional[int] = None vocab_size: Optional[int] = None
dim: Optional[int] = None hidden_size: Optional[int] = None
n_layers: Optional[int] = None num_hidden_layers: Optional[int] = None
norm_eps: Optional[float] = None rms_norm_eps: Optional[float] = None
dim_ffn: Optional[int] = None intermediate_size: Optional[int] = None
max_len: Optional[int] = None max_position_embeddings: Optional[int] = None
rope_theta: Optional[float] = None rope_theta: Optional[float] = None
rope_scaling: Optional[dict] = None rope_scaling: Optional[dict] = None
attn_type: str = "gqa" attn_type: str = "gqa"
n_heads: Optional[int] = None num_attention_heads: Optional[int] = None
n_kv_heads: Optional[int] = None num_key_value_heads: Optional[int] = None
use_qk_norm: Optional[bool] = None use_qk_norm: Optional[bool] = None
use_gated_attention: Optional[bool] = None use_gated_attention: Optional[bool] = None
+3
View File
@@ -45,6 +45,8 @@ class ProcessingConfig(BaseConfig):
Maximum number of characters to keep (default: 2_000_000). Maximum number of characters to keep (default: 2_000_000).
max_items : Optional[int] max_items : Optional[int]
Maximum number of items to process (default: None, unlimited). Maximum number of items to process (default: None, unlimited).
batch_size : int
Number of records tokenized together (default: 256).
packing_strategy : str packing_strategy : str
How to pack sequences into a contiguous stream. How to pack sequences into a contiguous stream.
@@ -65,6 +67,7 @@ class ProcessingConfig(BaseConfig):
min_chars: int = 50 min_chars: int = 50
max_chars: int = 2_000_000 max_chars: int = 2_000_000
max_items: Optional[int] = None max_items: Optional[int] = None
batch_size: int = 256
packing_strategy: str = "simple" packing_strategy: str = "simple"
max_packed_len: int = 8192 max_packed_len: int = 8192
truncation_mode: str = "keep_start" truncation_mode: str = "keep_start"
+33 -2
View File
@@ -37,8 +37,9 @@ class TrainConfig(BaseConfig):
grad_accum_steps: int = field( grad_accum_steps: int = field(
default=1, metadata={"help": "Number of iterations between steps."} default=1, metadata={"help": "Number of iterations between steps."}
) )
max_grad_norm: float = field( max_grad_norm: Optional[float] = field(
default=1.0, metadata={"help": "Maximum gradient norm."} default=1.0,
metadata={"help": "Maximum gradient norm. None disables clipping."},
) )
gradient_checkpointing_modules: List[str] = field( gradient_checkpointing_modules: List[str] = field(
default_factory=list, default_factory=list,
@@ -87,6 +88,10 @@ class TrainConfig(BaseConfig):
pin_memory: bool = field( pin_memory: bool = field(
default=False, metadata={"help": "Pin memory for dataloader."} default=False, metadata={"help": "Pin memory for dataloader."}
) )
collate_fn: Optional[Callable[[List[Any]], Any]] = field(
default=None,
metadata={"help": "Collate function for dataloader (e.g. dpo_collate_fn)."},
)
# distributed training # distributed training
nprocs: int = field( nprocs: int = field(
@@ -133,6 +138,32 @@ class TrainConfig(BaseConfig):
metadata={"help": "NEFTune noise alpha (0=disabled, typical: 5.0)."}, metadata={"help": "NEFTune noise alpha (0=disabled, typical: 5.0)."},
) )
# 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)."
},
)
executor_kwargs: Dict[str, Any] = field( executor_kwargs: Dict[str, Any] = field(
default_factory=dict, default_factory=dict,
metadata={"help": "Extra kwargs passed to ExecutorFactory.create()."}, metadata={"help": "Extra kwargs passed to ExecutorFactory.create()."},
+8 -2
View File
@@ -1,15 +1,18 @@
from astrai.dataset.dataset import ( from astrai.dataset.dataset import (
BaseDataset, BaseDataset,
DatasetFactory, DatasetFactory,
dpo_collate_fn,
grpo_collate_fn, grpo_collate_fn,
) )
from astrai.dataset.sampler import ResumableDistributedSampler from astrai.dataset.sampler import RDSampler
from astrai.dataset.storage import ( from astrai.dataset.storage import (
H5Store, H5Store,
JsonlStore, JsonlStore,
MmapStore, MmapStore,
Recordable,
Store, Store,
StoreFactory, StoreFactory,
Streamable,
detect_format, detect_format,
) )
from astrai.serialization import ( from astrai.serialization import (
@@ -22,8 +25,11 @@ from astrai.serialization import (
__all__ = [ __all__ = [
"BaseDataset", "BaseDataset",
"DatasetFactory", "DatasetFactory",
"dpo_collate_fn",
"grpo_collate_fn", "grpo_collate_fn",
"Store", "Store",
"Streamable",
"Recordable",
"StoreFactory", "StoreFactory",
"H5Store", "H5Store",
"MmapStore", "MmapStore",
@@ -33,5 +39,5 @@ __all__ = [
"load_h5", "load_h5",
"save_bin", "save_bin",
"load_bin", "load_bin",
"ResumableDistributedSampler", "RDSampler",
] ]
+359 -230
View File
@@ -1,7 +1,31 @@
"""Dataset implementations with factory pattern for training.""" """Dataset implementations for training.
Composition over inheritance — every dataset is a thin wrapper that
binds a :class:`Store` to a particular train-type's key mapping. All
sample-id → token/record indexing lives on the Store; datasets never
know about window/stride math or segment layouts.
Class hierarchy:
BaseDataset (ABC) — holds a Store, exposes __len__/keys,
overrides __getitem__
├── SEQDataset — next-token prediction (stream)
├── SFTDataset — loss-mask + position_ids (stream)
├── DPODataset — chosen/rejected pairs (record)
└── GRPODataset — prompt + response group (record)
``DatasetFactory.load(train_type, load_path, window_size, stride, …)``
builds the Store (auto-detecting format) before constructing the
matching dataset. Passing ``store=`` skips Store construction.
When a record dataset (DPO) reads from raw JSONL, a *processor*
function (pure ``record -> Dict[str, Tensor]``) is forwarded to
:class:`JsonlStore` so tokenisation happens on the fly.
"""
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Dict, List, Optional from functools import partial
from typing import Callable, Dict, List, Optional
import torch import torch
from torch import Tensor from torch import Tensor
@@ -13,6 +37,147 @@ from astrai.dataset.storage import (
detect_format, detect_format,
) )
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.tokenize import AutoTokenizer
def dpo_tokenize(
record: dict,
tokenizer,
max_len: int = 2048,
) -> Optional[dict]:
"""Tokenize one DPO record into chosen/rejected + masks.
Applies the tokenizer's chat template so token sequences match the
SFT checkpoint's format. Prompt is rendered with
``add_generation_prompt=True``; chosen/rejected are appended as a
single assistant turn.
Accepts:
- Flat: ``{"prompt": str, "chosen": str, "rejected": str}``
- Conv: ``{"prompt": [{role, content}, ...], "chosen": [...], ...}``
- Legacy: ``{"input": str, "chosen": str, "rejected": str}``
No packing, no ``position_ids`` — DPO sequences are independent.
"""
prompt = record.get("prompt") or record.get("input")
chosen = record.get("chosen")
rejected = record.get("rejected")
if prompt is None or chosen is None or rejected is None:
return None
prompt_messages = _to_messages(prompt)
chosen_text = _extract_text(chosen)
rejected_text = _extract_text(rejected)
if chosen_text is None or rejected_text is None:
return None
chosen_messages = prompt_messages + [{"role": "assistant", "content": chosen_text}]
rejected_messages = prompt_messages + [
{"role": "assistant", "content": rejected_text}
]
prompt_ids = tokenizer.apply_chat_template(
prompt_messages, tokenize=True, add_generation_prompt=True
)
ch_ids = tokenizer.apply_chat_template(
chosen_messages, tokenize=True, add_generation_prompt=False
)
re_ids = tokenizer.apply_chat_template(
rejected_messages, tokenize=True, add_generation_prompt=False
)
full_ch = ch_ids[:max_len]
full_re = re_ids[:max_len]
prompt_len = min(len(prompt_ids), max_len)
ch_mask = [0] * prompt_len + [1] * max(0, len(full_ch) - prompt_len)
ch_mask = ch_mask[:max_len]
re_mask = [0] * prompt_len + [1] * max(0, len(full_re) - prompt_len)
re_mask = re_mask[:max_len]
return {
"chosen": full_ch,
"rejected": full_re,
"chosen_mask": ch_mask,
"rejected_mask": re_mask,
}
def _to_messages(value) -> list:
"""Accept str or conversation list; return message list."""
if isinstance(value, str):
return [{"role": "user", "content": value}]
if isinstance(value, list):
return value
return [{"role": "user", "content": str(value)}]
def _extract_text(value) -> Optional[str]:
"""Accept str or conversation list; return plain text."""
if value is None:
return None
if isinstance(value, str):
return value
if isinstance(value, list):
return "".join(m.get("content", "") for m in value if isinstance(m, dict))
return None
def dpo_processor(
record: dict,
tokenizer,
max_len: int = 2048,
) -> Dict[str, Tensor]:
"""DPO processor: wraps :func:`dpo_tokenize` and returns tensors."""
result = dpo_tokenize(record, tokenizer, max_len=max_len)
if result is None:
raise ValueError(f"Malformed DPO record: {list(record.keys())}")
return {
"chosen": torch.tensor(result["chosen"], dtype=torch.int32),
"rejected": torch.tensor(result["rejected"], dtype=torch.int32),
"chosen_mask": torch.tensor(result["chosen_mask"], dtype=torch.bool),
"rejected_mask": torch.tensor(result["rejected_mask"], dtype=torch.bool),
}
def dpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
"""Collate variable-length DPO samples into padded 2-D tensors.
Input: list of dicts, each with:
- chosen: [C_i]
- rejected: [R_i]
- chosen_mask: [C_i]
- rejected_mask: [R_i]
Output (padded to the max length across chosen/rejected within the batch):
- chosen: [B, S_max]
- rejected: [B, S_max]
- chosen_mask: [B, S_max]
- rejected_mask: [B, S_max]
"""
B = len(batch)
S_max = max(b["chosen"].size(0) for b in batch)
S_max = max(S_max, max(b["rejected"].size(0) for b in batch))
chosen = torch.zeros(B, S_max, dtype=torch.long)
rejected = torch.zeros(B, S_max, dtype=torch.long)
chosen_mask = torch.zeros(B, S_max, dtype=torch.bool)
rejected_mask = torch.zeros(B, S_max, dtype=torch.bool)
for i, b in enumerate(batch):
c_len = b["chosen"].size(0)
r_len = b["rejected"].size(0)
chosen[i, :c_len] = b["chosen"]
rejected[i, :r_len] = b["rejected"]
chosen_mask[i, :c_len] = b["chosen_mask"]
rejected_mask[i, :r_len] = b["rejected_mask"]
return {
"chosen": chosen,
"rejected": rejected,
"chosen_mask": chosen_mask,
"rejected_mask": rejected_mask,
}
def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]: def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
@@ -25,7 +190,8 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
- rewards: [G] - rewards: [G]
Output: Output:
- prompts: [B, P_max] - prompts: [B, P_max], left-padded
- prompt_mask: [B, P_max]
- responses: [B, G, R_max] - responses: [B, G, R_max]
- masks: [B, G, R_max] - masks: [B, G, R_max]
- rewards: [B, G] - rewards: [B, G]
@@ -36,13 +202,15 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
R_max = max(r.size(0) for b in batch for r in b["responses"]) R_max = max(r.size(0) for b in batch for r in b["responses"])
prompts = torch.zeros(B, P_max, dtype=torch.long) prompts = torch.zeros(B, P_max, dtype=torch.long)
prompt_mask = torch.zeros(B, P_max, dtype=torch.bool)
responses = torch.zeros(B, G, R_max, dtype=torch.long) responses = torch.zeros(B, G, R_max, dtype=torch.long)
masks = torch.zeros(B, G, R_max, dtype=torch.bool) masks = torch.zeros(B, G, R_max, dtype=torch.bool)
rewards = torch.zeros(B, G, dtype=torch.float32) rewards = torch.zeros(B, G, dtype=torch.float32)
for i, b in enumerate(batch): for i, b in enumerate(batch):
p_len = b["prompts"].size(0) p_len = b["prompts"].size(0)
prompts[i, :p_len] = b["prompts"] prompts[i, -p_len:] = b["prompts"]
prompt_mask[i, -p_len:] = True
rewards[i, : b["rewards"].size(0)] = b["rewards"] rewards[i, : b["rewards"].size(0)] = b["rewards"]
for g in range(min(G, len(b["responses"]))): for g in range(min(G, len(b["responses"]))):
r_len = b["responses"][g].size(0) r_len = b["responses"][g].size(0)
@@ -52,206 +220,218 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
return { return {
"prompts": prompts, "prompts": prompts,
"prompt_mask": prompt_mask,
"responses": responses, "responses": responses,
"masks": masks, "masks": masks,
"rewards": rewards, "rewards": rewards,
} }
class BaseDataset(Dataset, ABC): def validate_keys(store: Store, required: List[str]) -> None:
"""Abstract base class for all dataset types. """Raise ``KeyError`` if *store* is missing any *required* key."""
if not required:
Implements common functionality for window-based data fetching.
Uses a storage abstraction for format-agnostic data loading.
"""
def __init__(self, window_size: int, stride: int):
super().__init__()
self.window_size = window_size
self.stride = stride
self.storage: Optional[Store] = None
@property
def required_keys(self) -> List[str]:
"""Return required storage keys for this dataset type.
Subclasses should override to specify expected keys.
"""
return []
def _validate_keys(self):
if not self.required_keys:
return return
actual_keys = set(self.storage.keys) actual = set(store.keys)
missing = [k for k in self.required_keys if k not in actual_keys] missing = [k for k in required if k not in actual]
if missing: if missing:
raise KeyError( raise KeyError(
f"Dataset {type(self).__name__} requires keys {self.required_keys}, " f"Store at {getattr(store, '_load_path', '?')} is missing required "
f"but storage at {self._load_path} only has {sorted(actual_keys)}. " f"keys {missing}; available keys are {sorted(actual)}."
f"Missing: {missing}"
) )
def load(self, load_path: str, storage_type: Optional[str] = None, **kwargs):
"""Load dataset from the given path.
Auto-detects the storage format if not specified. class BaseDataset(Dataset, ABC):
"""Abstract base class for dataset types.
Args: Holds a :class:`Store`. All sample-id indexing is delegated to the
load_path: Path to the data directory or file store — this class exposes ``__len__`` as ``len(store)`` and the
storage_type: Force a specific storage type ("h5", "bin", "jsonl"), ``keys`` property as ``store.keys``. Subclasses implement
or None for auto-detection ``__getitem__`` with the train-type-specific key mapping and any
**kwargs: Extra arguments forwarded to the store constructor and training-only index arithmetic (e.g. the next-token ``+1`` shift).
to ``store.load()``.
Raises:
KeyError: If the loaded storage is missing required keys.
""" """
if storage_type is None:
storage_type = detect_format(load_path)
self.storage = StoreFactory.create(storage_type, **kwargs)
self._load_path = load_path
self.storage.load(load_path, **kwargs)
self._validate_keys()
@property required_keys: List[str] = []
def count(self) -> int:
"""Return the total number of raw elements (tokens) in the dataset.""" def __init__(self, store: Store):
if self.storage is None: super().__init__()
return 0 self.store: Store = store
return len(self.storage) validate_keys(store, self.required_keys)
def __len__(self) -> int:
return len(self.store)
@property @property
def keys(self) -> List[str]: def keys(self) -> List[str]:
"""Return the available data keys.""" return self.store.keys
if self.storage is None:
return []
return self.storage.keys
def get_index(self, index: int) -> tuple: @property
"""Calculate begin and end indices for a sample. def token_count(self) -> int:
return self.store.token_count
Args:
index: Sample index
Returns:
Tuple of (begin_idx, end_idx)
"""
if self.storage is None:
raise RuntimeError("Dataset not loaded, call load() first")
total = len(self.storage)
if total <= self.window_size:
raise ValueError(
f"Data too short: {total} tokens <= window_size {self.window_size}"
)
begin_idx = min(index * self.stride, total - 1 - self.window_size)
end_idx = min(begin_idx + self.window_size, total - 1)
return begin_idx, end_idx
@abstractmethod @abstractmethod
def __getitem__(self, index: int) -> Dict[str, Tensor]: def __getitem__(self, index: int) -> Dict[str, Tensor]:
"""Get a single sample by index.
Must be implemented by subclasses.
"""
raise NotImplementedError raise NotImplementedError
def __len__(self) -> int:
if self.storage is None:
return 0
total = len(self.storage)
if total <= self.window_size:
return 0
return (total - 1 - self.window_size) // self.stride + 1
class DatasetFactory(BaseFactory["BaseDataset"]): class DatasetFactory(BaseFactory["BaseDataset"]):
"""Factory class for creating dataset instances. """Factory for creating dataset instances by train-type.
Supports decorator-based registration for extensible dataset types. Use :meth:`DatasetFactory.register("custom")` to register new
All default dataset types (seq, sft, dpo, grpo) are registered automatically dataset classes; they must inherit from :class:`BaseDataset`.
when their classes are defined with the decorator.
Example usage:
@DatasetFactory.register("custom")
class CustomDataset(BaseDataset):
...
dataset = DatasetFactory.create("custom", window_size, stride)
""" """
@classmethod @classmethod
def load( def load(
cls, cls,
train_type: str, train_type: str,
load_path: str, load_path: Optional[str] = None,
window_size: int, window_size: int = 0,
stride: Optional[int] = None, stride: Optional[int] = None,
storage_type: Optional[str] = None, storage_type: Optional[str] = None,
tokenizer_path: Optional[str] = None,
max_len: int = 2048,
store: Optional[Store] = None,
**kwargs, **kwargs,
) -> "BaseDataset": ) -> "BaseDataset":
"""Create and load a dataset in one step. """Create and load a dataset in one step.
Two entry points:
- **store given**: bind it directly — the caller fully controls
Store construction and processor setup. *load_path*,
*storage_type*, *tokenizer_path*, *window_size*, *stride* are
ignored.
- **store is None**: build a Store from *load_path*, auto-detecting
format and constructing a processor when *tokenizer_path* is
given for a record dataset on JSONL.
Args: Args:
train_type: Type of training dataset train_type: Registered dataset name ("seq", "sft", "dpo",
load_path: Path to the data file "grpo", …).
window_size: Window size for data sampling load_path: Path to the data file or directory (ignored if
stride: Stride between consecutive samples (default: same as window_size) *store* is given).
storage_type: Storage type ("h5", "bin", "jsonl") or None for auto-detection window_size: Stream window length — only meaningful for
**kwargs: Extra arguments forwarded to ``dataset.load()``. 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
None for auto-detection.
tokenizer_path: Path to tokenizer for lazy JSONL
tokenisation (record datasets only).
max_len: Max sequence length forwarded to processors.
store: Pre-built, already-loaded Store instance.
**kwargs: Extra arguments forwarded to ``store.load()``.
Returns: Returns:
Loaded dataset instance Loaded dataset instance.
""" """
if store is not None:
return cls.create(train_type, store=store)
if load_path is None:
raise ValueError("Either load_path or store must be provided")
if storage_type is None:
storage_type = detect_format(load_path)
if stride is None: if stride is None:
stride = window_size stride = window_size
dataset = cls.create(train_type, window_size, stride) processor = cls._maybe_build_processor(
dataset.load(load_path, storage_type=storage_type, **kwargs) train_type, storage_type, tokenizer_path, max_len
)
return dataset store_window = cls._store_window_for(train_type, window_size)
store = StoreFactory.create(
storage_type,
window_size=store_window,
stride=stride if stride else store_window,
)
if processor is not None:
store.load(load_path, processor=processor, **kwargs)
else:
load_kwargs = dict(kwargs)
if (
tokenizer_path is not None
and storage_type == "jsonl"
and train_type in ("seq", "sft")
and "tokenizer_path" not in load_kwargs
):
load_kwargs["tokenizer_path"] = tokenizer_path
store.load(load_path, **load_kwargs)
return cls.create(train_type, store=store)
@staticmethod
def _store_window_for(train_type: str, window_size: int) -> int:
"""Stream datasets consume ``window_size``; record datasets ignore it.
Record datasets (dpo/grpo) treat each record as an independent
training unit and never window, so the store is built with
``window_size=0`` and ``len(store)`` returns the record count.
"""
if train_type in ("seq", "sft"):
return window_size
return 0
@staticmethod
def _maybe_build_processor(
train_type: str,
storage_type: str,
tokenizer_path: Optional[str],
max_len: int,
) -> Optional[Callable[[dict], Dict[str, Tensor]]]:
"""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)
return ``None`` so no tokenizer is loaded.
"""
if tokenizer_path is None or storage_type != "jsonl":
return None
if train_type == "dpo":
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
return None
@DatasetFactory.register("seq") @DatasetFactory.register("seq")
class SEQDataset(BaseDataset): class SEQDataset(BaseDataset):
"""Dataset for sequential next-token prediction training.""" """Dataset for sequential next-token prediction training.
@property Stream mode: ``store.fetch(begin, end, "sequence")`` returns the
def required_keys(self) -> List[str]: input window; the +1 shifted call returns the next-token target.
return ["sequence"] """
def _fetch_data(self, begin_idx: int, end_idx: int) -> Tensor: required_keys = ["sequence"]
return self.storage.fetch(begin_idx, end_idx, "sequence")
def __getitem__(self, index): def __getitem__(self, index: int):
begin_idx, end_idx = self.get_index(index) begin, end = self.store.sample_window(index)
x = self.store.fetch(begin, end, "sequence")
x = self._fetch_data(begin_idx, end_idx).to(dtype=torch.long) y = self.store.fetch(begin + 1, end + 1, "sequence")
y = self._fetch_data(begin_idx + 1, end_idx + 1).to(dtype=torch.long) return {
"input_ids": x.to(dtype=torch.long),
return {"input_ids": x, "target_ids": y} "target_ids": y.to(dtype=torch.long),
}
@DatasetFactory.register("sft") @DatasetFactory.register("sft")
class SFTDataset(BaseDataset): class SFTDataset(BaseDataset):
"""Dataset for supervised fine-tuning with loss masking.""" """Dataset for supervised fine-tuning with loss masking.
@property Stream mode: ``sequence``/``loss_mask``/``position_ids`` are sliced
def required_keys(self) -> List[str]: to the window. ``loss_mask`` and ``target_ids`` use the +1 shifted
return ["sequence", "loss_mask", "position_ids"] slice so they align with the predicted positions.
"""
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor: required_keys = ["sequence", "loss_mask", "position_ids"]
return self.storage.fetch(begin_idx, end_idx, key)
def __getitem__(self, index):
begin_idx, end_idx = self.get_index(index)
x = self._fetch_data(begin_idx, end_idx, "sequence")
y = self._fetch_data(begin_idx + 1, end_idx + 1, "sequence")
position_ids = self._fetch_data(begin_idx, end_idx, "position_ids")
loss_mask = self._fetch_data(begin_idx + 1, end_idx + 1, "loss_mask")
def __getitem__(self, index: int):
begin, end = self.store.sample_window(index)
x = self.store.fetch(begin, end, "sequence")
y = self.store.fetch(begin + 1, end + 1, "sequence")
position_ids = self.store.fetch(begin, end, "position_ids")
loss_mask = self.store.fetch(begin + 1, end + 1, "loss_mask")
return { return {
"input_ids": x.to(dtype=torch.long), "input_ids": x.to(dtype=torch.long),
"target_ids": y.to(dtype=torch.long), "target_ids": y.to(dtype=torch.long),
@@ -262,32 +442,37 @@ class SFTDataset(BaseDataset):
@DatasetFactory.register("dpo") @DatasetFactory.register("dpo")
class DPODataset(BaseDataset): class DPODataset(BaseDataset):
"""Dataset for Direct Preference Optimization training.""" """Record-structured dataset for Direct Preference Optimization.
@property Each sample is one preference pair (chosen + rejected) and is an
def required_keys(self) -> List[str]: independent training unit — no windowing, stride, or cross-record
return ["chosen", "rejected", "chosen_mask", "rejected_mask"] concatenation. This keeps each sequence self-contained so attention
never leaks across preference pairs.
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor: Two loading paths (handled by :class:`DatasetFactory`):
return self.storage.fetch(begin_idx, end_idx, key)
def __getitem__(self, index: int): - **Pre-tokenized** (H5/bin): ``store.load(path)`` reads per-record
begin_idx, end_idx = self.get_index(index) tensors; ``__getitem__`` returns them directly.
- **Raw JSONL** (``tokenizer_path=...``): builds a lazy processor
via :func:`dpo_processor` that tokenises on the fly — no packing,
no ``position_ids``.
"""
chosen = self._fetch_data(begin_idx, end_idx, "chosen").to(dtype=torch.long) required_keys = ["chosen", "rejected", "chosen_mask", "rejected_mask"]
rejected = self._fetch_data(begin_idx, end_idx, "rejected").to(dtype=torch.long)
chosen_mask = self._fetch_data(begin_idx, end_idx, "chosen_mask").to(
dtype=torch.bool
)
rejected_mask = self._fetch_data(begin_idx, end_idx, "rejected_mask").to(
dtype=torch.bool
)
def make_processor(self, tokenizer, max_len: int):
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
def __getitem__(self, index: int) -> Dict[str, Tensor]:
return { return {
"chosen": chosen, "chosen": self.store.fetch_record(index, "chosen").to(dtype=torch.long),
"rejected": rejected, "rejected": self.store.fetch_record(index, "rejected").to(dtype=torch.long),
"chosen_mask": chosen_mask, "chosen_mask": self.store.fetch_record(index, "chosen_mask").to(
"rejected_mask": rejected_mask, dtype=torch.bool
),
"rejected_mask": self.store.fetch_record(index, "rejected_mask").to(
dtype=torch.bool
),
} }
@@ -295,10 +480,8 @@ class DPODataset(BaseDataset):
class GRPODataset(BaseDataset): class GRPODataset(BaseDataset):
"""Dataset for offline Group Relative Policy Optimization. """Dataset for offline Group Relative Policy Optimization.
Unlike the window-based datasets (SEQ/SFT/DPO), GRPO data is Each sample is one prompt with its group of responses and scalar
record-structured: each sample is one prompt with its group of rewards — an independent training unit with no windowing or stride.
responses and scalar rewards. There is no windowing or stride —
every record is an independent training unit.
Expected storage layout (produced by JsonlStore or pre-tokenized): Expected storage layout (produced by JsonlStore or pre-tokenized):
@@ -308,70 +491,16 @@ class GRPODataset(BaseDataset):
- ``rewards``: List[Tensor] — one 1-D float tensor (len G) per record - ``rewards``: List[Tensor] — one 1-D float tensor (len G) per record
""" """
def __init__(self, window_size: int = 0, stride: int = 0, **kwargs): required_keys = ["prompts", "responses", "masks", "rewards"]
super().__init__(window_size=window_size, stride=stride or window_size)
self._records: List[dict] = []
@property
def required_keys(self) -> List[str]:
return ["prompts", "responses", "masks", "rewards"]
def load(self, load_path: str, storage_type: Optional[str] = None, **kwargs):
if storage_type is None:
storage_type = detect_format(load_path)
self.storage = StoreFactory.create(storage_type, **kwargs)
self._load_path = load_path
self.storage.load(load_path, **kwargs)
self._validate_keys()
self._build_records()
def _validate_keys(self):
actual_keys = set(self.storage.keys)
missing = [k for k in self.required_keys if k not in actual_keys]
if missing:
raise KeyError(
f"GRPODataset requires keys {self.required_keys}, "
f"but storage only has {sorted(actual_keys)}. Missing: {missing}"
)
def _build_records(self):
"""Unfold segmented storage into per-record lists.
``prompts`` is a flat list of 1-D tensors (one per record).
``responses`` / ``masks`` are nested lists (G tensors per record).
``rewards`` is a flat list of 1-D tensors (len G per record).
"""
prompt_segs = self.storage._data.get("prompts", [])
response_segs = self.storage._data.get("responses", [])
mask_segs = self.storage._data.get("masks", [])
reward_segs = self.storage._data.get("rewards", [])
n_records = len(prompt_segs)
self._records = []
for i in range(n_records):
self._records.append(
{
"prompts": prompt_segs[i],
"responses": response_segs[i] if i < len(response_segs) else [],
"masks": mask_segs[i] if i < len(mask_segs) else [],
"rewards": reward_segs[i]
if i < len(reward_segs)
else torch.tensor([]),
}
)
@property
def count(self) -> int:
return len(self._records)
def __len__(self) -> int:
return len(self._records)
def __getitem__(self, index: int) -> Dict[str, Tensor]: def __getitem__(self, index: int) -> Dict[str, Tensor]:
rec = self._records[index] prompts = self.store.fetch_record(index, "prompts")
responses = self.store.fetch_record(index, "responses")
masks = self.store.fetch_record(index, "masks")
rewards = self.store.fetch_record(index, "rewards")
return { return {
"prompts": rec["prompts"].to(dtype=torch.long), "prompts": prompts.to(dtype=torch.long),
"responses": [r.to(dtype=torch.long) for r in rec["responses"]], "responses": [r.to(dtype=torch.long) for r in responses],
"masks": [m.to(dtype=torch.bool) for m in rec["masks"]], "masks": [m.to(dtype=torch.bool) for m in masks],
"rewards": rec["rewards"].to(dtype=torch.float32), "rewards": rewards.to(dtype=torch.float32),
} }
+9 -1
View File
@@ -5,7 +5,15 @@ import torch.distributed as dist
from torch.utils.data import Dataset, Sampler from torch.utils.data import Dataset, Sampler
class ResumableDistributedSampler(Sampler[int]): class RDSampler(Sampler[int]):
"""Resumable Distributed Sampler.
A distributed sampler that supports checkpoint-based resume: iteration
state (epoch, position) is tracked so training can continue from the
exact sample after a restart. Shards the dataset across
``dist.world_size`` replicas with optional shuffling.
"""
def __init__( def __init__(
self, self,
data_source: Dataset, data_source: Dataset,
+493 -174
View File
@@ -1,20 +1,48 @@
"""Storage backends for different data formats. """Storage backends for different data formats.
Layers: Architecture (composition over inheritance):
- I/O layer: save_* / load_* functions, read/write raw files (HDF5/bin)
return Dict[str, List[Tensor]] — format-specific, no state
- Store (ABC): central abstraction, normalizes multi-segment into
Dict[str, List[Tensor]] per key via _normalize(),
fetch() uses bisect across segments — no forced concat
- Dataset layer: BaseDataset owns a Store, only calls store.fetch(begin, end, key)
Key properties: Store (ABC) — owns _data/_cum/_offsets bookkeeping
- Multi-segment: segments kept as-is, no forced concatenation — safe for + window_size/stride for sample-id
datasets larger than RAM indexing. __getitem__/__len__ produce
- Explicit length: _length = min(total elements across keys), set at load, the smallest iterable unit so Dataset
__len__ returns O(1) classes are pure delegators.
- Zero-copy mmap: MmapStore wraps np.memmap(mode="r"), all DataLoader Streamable (mixin) — raw token slice fetch(begin, end, keys)
workers share OS page-cache pages Recordable (mixin) — raw record slice fetch_record(idx, keys)
H5Store(Store, Streamable, Recordable)
MmapStore(Store, Streamable, Recordable)
JsonlStore(Store, Streamable, Recordable)
Each mixin is a stateless trait that relies on ``self._data`` etc.
provided by :class:`Store`. Concrete stores mix in whichever access
primitives they support — ``Store`` is the sole base class, so there is
no diamond inheritance or MRO ambiguity.
Sample-id indexing lives on :class:`Store`, not on the dataset:
- **Stream mode** (``window_size > 0``): ``len(store)`` returns the number
of ``(window_size, stride)`` windows that fit in the token river;
``store[i]`` returns the *i*-th window as a dict of per-key tensors;
``store.sample_window(i)`` exposes the underlying ``(begin, end)``
token slice for callers (e.g. next-token trainers) that need a +1
shifted companion window.
- **Record mode** (``num_records > 0``): ``len(store)`` returns the
record count; ``store[i]`` returns the *i*-th record dict.
Raw token/record access via :meth:`fetch` / :meth:`fetch_record`
remains available for low-level callers that want explicit index
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.
:class:`JsonlStore` supports a lazy mode (``processor=fn``) that keeps
raw records and defers tokenisation to ``fetch_record`` — used by DPO
to train directly from a ``.jsonl`` file without a pre-tokenised copy.
""" """
import bisect import bisect
@@ -23,20 +51,19 @@ import json
import logging import logging
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from pathlib import Path from pathlib import Path
from typing import Dict, List, Union from typing import Callable, Dict, List, Optional, Tuple, Union
import torch import torch
from torch import Tensor from torch import Tensor
from astrai.config.preprocess_config import PipelineConfig from astrai.config.preprocess_config import PipelineConfig
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.preprocessing.builder import MaskBuilderFactory from astrai.preprocessing.transform import TokenizeTransform
from astrai.preprocessing.position_id import PositionIdStrategyFactory
from astrai.serialization import ( from astrai.serialization import (
load_bin, load_bin,
load_bin_offsets,
load_h5, load_h5,
) )
from astrai.tokenize import AutoTokenizer
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -48,7 +75,7 @@ def detect_format(load_path: str) -> str:
load_path: Directory or file path load_path: Directory or file path
Returns: Returns:
Format string ("h5", "bin", or "jsonl") Format string ("h5", "bin", "jsonl", or "processed")
Raises: Raises:
FileNotFoundError: If no supported data files are found FileNotFoundError: If no supported data files are found
@@ -81,59 +108,254 @@ def detect_format(load_path: str) -> str:
] ]
if jsonl_files: if jsonl_files:
return "jsonl" return "jsonl"
json_files = [
Path(p) for p in glob.glob(str(root / "**" / "*.json"), recursive=True)
]
if json_files:
return "jsonl"
raise FileNotFoundError(f"No supported data files found at {load_path}") raise FileNotFoundError(f"No supported data files found at {load_path}")
class Store(ABC): class Store(ABC):
"""String keys -> segmented tensors with ``fetch(begin, end, keys)``. """Common base for all storage backends.
Each key maps to one or more tensor segments (no forced concatenation). A Store owns both its data layout AND its sample-id → token/record
``len(store)`` returns ``self._length`` (explicit, O(1)), the minimum index translation. Datasets are thin wrappers that bind a Store
total element count across all keys. to a particular train-type's key mapping; they never know about
window/stride math.
Subclasses fill ``self._data`` and ``self._cum`` during ``load()`` Two iteration modes:
via ``_normalize()``.
- **Stream** (``window_size > 0``): data is treated as one long
token river. ``len(store)`` returns the number of windows;
``store[i]`` slices every stream-compatible key to window ``i``;
``store.sample_window(i)`` returns the ``(begin, end)`` token
slice for callers needing a +1 shifted companion window.
- **Record** (``num_records > 0``): data is per-record.
``len(store)`` returns ``num_records``; ``store[i]`` returns
the *i*-th record as a dict.
Raw token slicing is still available via :meth:`fetch` (mixed in
by :class:`Streamable`) when a store has stream support configured.
Raw record slicing via :meth:`fetch_record` (mixed in by
:class:`Recordable`) when a store has record support.
``token_count`` exposes the raw total stream length — this is what
``len(store)`` returned in the legacy stream-only API and what
stream-bound ``fetch`` uses for its bounds check.
""" """
def __init__(self): segments_are_records: bool = False
def __init__(
self,
window_size: int = 0,
stride: Optional[int] = None,
):
self._data: Dict[str, List[Tensor]] = {} self._data: Dict[str, List[Tensor]] = {}
self._cum: Dict[str, List[int]] = {} self._cum: Dict[str, List[int]] = {}
self._offsets: Dict[str, List[int]] = {}
self._length: int = 0 self._length: int = 0
self._num_records: int = 0
self._window_size: int = int(window_size)
self._stride: int = int(stride) if stride is not None else int(window_size)
@abstractmethod @abstractmethod
def load(self, path: str) -> None: def load(self, path: str, **kwargs) -> None:
raise NotImplementedError raise NotImplementedError
@property @property
def keys(self) -> List[str]: def keys(self) -> List[str]:
return list(self._data.keys()) return list(self._data.keys())
def __len__(self) -> int: @property
def window_size(self) -> int:
return self._window_size
@property
def stride(self) -> int:
return self._stride
@property
def token_count(self) -> int:
"""Total tokens across all stream segments.
Useful for the bounds-checked raw :meth:`fetch` and as the
legacy ``len(store)`` value.
"""
return self._length return self._length
@property
def num_records(self) -> int:
"""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``).
"""
return self._num_records
@property
def num_samples(self) -> int:
"""Number of items produced by ``__getitem__``.
Stream-mode wins when ``window_size > 0`` and there are tokens
to slice; otherwise falls back to ``num_records``.
"""
if self._window_size > 0 and self._length > 0:
total = self._length
w = self._window_size
if total <= w:
return 0
return (total - 1 - w) // self._stride + 1
return self._num_records
def __len__(self) -> int:
return self.num_samples
def __getitem__(self, index: int) -> Dict[str, Tensor]:
if index < 0:
index += self.num_samples
if not 0 <= index < self.num_samples:
raise IndexError(
f"Store index out of range: {index}, num_samples={self.num_samples}"
)
if self._window_size > 0 and self._length > 0:
begin, end = self.sample_window(index)
keys = self._stream_keys()
return {k: self.fetch(begin, end, k) for k in keys}
return self.fetch_record(index, self._record_keys())
def sample_window(self, index: int) -> Tuple[int, int]:
"""Return ``(begin, end)`` token positions for stream sample *index*.
The clipped tail keeps the last reachable window inside the
token river instead of overshooting. Caller is responsible
for staying within :attr:`num_samples`: an out-of-range index
raises ``IndexError``.
"""
if self._window_size <= 0:
raise RuntimeError("sample_window() requires window_size > 0 (stream mode)")
if self._window_size <= 0 or self._length <= self._window_size:
raise IndexError(
f"Data too short for window: token_count={self._length}, "
f"window_size={self._window_size}"
)
if not 0 <= index < self.num_samples:
raise IndexError(
f"Sample index out of range: {index}, num_samples={self.num_samples}"
)
total = self._length
begin = min(index * self._stride, total - 1 - self._window_size)
end = min(begin + self._window_size, total - 1)
return begin, end
def _stream_keys(self) -> List[str]:
out: List[str] = []
for k, tensors in self._data.items():
if tensors and isinstance(tensors[0], list):
continue
out.append(k)
return out
def _record_keys(self) -> List[str]:
return list(self._data.keys())
def _normalize(
self,
raw: Dict[str, list],
offsets: Optional[Dict[str, List[int]]] = None,
):
"""Register segments and pre-compute indices for both access modes.
Stream mode: ``_cum[key]`` accumulates per-segment lengths so
``Streamable._fetch_stream_key`` can bisect across segments
without concatenation.
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
a per-record list and ``fetch_record`` indexes it directly.
Nested keys (GRPO ``responses``/``masks`` as
``List[List[Tensor]]``) are stored as-is and excluded from both
cumulative bookkeepings — they are only accessed record-by-record.
"""
flat_lengths = []
for key, tensors in raw.items():
self._data[key] = tensors
if not tensors:
self._cum[key] = []
flat_lengths.append(0)
continue
if isinstance(tensors[0], list):
self._cum[key] = []
continue
cum = []
total = 0
for t in tensors:
total += t.shape[0]
cum.append(total)
self._cum[key] = cum
flat_lengths.append(cum[-1] if cum else 0)
self._length = min(flat_lengths) if flat_lengths else 0
valid_offsets: Dict[str, List[int]] = {}
if offsets:
for key, off in offsets.items():
segs = self._data.get(key, [])
if len(segs) == 1 and len(off) > 1:
valid_offsets[key] = off
elif len(segs) > 1:
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.",
key,
len(segs),
)
self._offsets = valid_offsets
if valid_offsets:
record_counts = [len(v) - 1 for v in valid_offsets.values()]
self._num_records = min(record_counts) if record_counts else 0
elif self.segments_are_records:
per_record_counts = []
for key, tensors in self._data.items():
if tensors and isinstance(tensors[0], list):
continue
per_record_counts.append(len(tensors))
self._num_records = min(per_record_counts) if per_record_counts else 0
else:
self._num_records = 0
class Streamable:
"""Mixin granting raw token-stream access via :meth:`fetch`.
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
``fetch_record`` API from :class:`Recordable` is used instead.
"""
def fetch( def fetch(
self, self,
begin: int, begin: int,
end: int, end: int,
keys: Union[str, List[str]], keys: Union[str, List[str]],
): ):
if not self._data: return _stream_fetch(self, begin, end, keys)
def _stream_fetch(self, begin: int, end: int, keys: Union[str, List[str]]):
if not getattr(self, "_data", None):
raise RuntimeError("Store not loaded") raise RuntimeError("Store not loaded")
if not (0 <= begin < self._length and 0 <= end <= self._length): if not (0 <= begin < self._length and 0 <= end <= self._length):
raise ValueError( raise ValueError(
f"Index out of bounds: begin={begin}, end={end}, length={self._length}" f"Index out of bounds: begin={begin}, end={end}, length={self._length}"
) )
if isinstance(keys, str): if isinstance(keys, str):
return self._fetch_key(keys, begin, end) return _fetch_stream_key(self, keys, begin, end)
return {k: self._fetch_key(k, begin, end) for k in keys} return {k: _fetch_stream_key(self, k, begin, end) for k in keys}
def _fetch_key(self, key: str, begin: int, end: int) -> Tensor:
"""Fetch slice [begin, end) across potentially multiple segments.""" def _fetch_stream_key(self, key: str, begin: int, end: int) -> Tensor:
segments = self._data[key] segments = self._data[key]
cum = self._cum[key] cum = self._cum[key]
seg_start = bisect.bisect_right(cum, begin) seg_start = bisect.bisect_right(cum, begin)
@@ -148,198 +370,295 @@ class Store(ABC):
return results[0] if len(results) == 1 else torch.cat(results, dim=0) return results[0] if len(results) == 1 else torch.cat(results, dim=0)
def _normalize(self, raw: Dict[str, list]):
"""Register segments and pre-compute cumulative lengths.
Does NOT concatenate — segments are kept as-is to avoid OOM on class Recordable:
large datasets. Sets ``self._length`` to the minimum total """Mixin granting raw record access via :meth:`fetch_record`.
element count across all flat-tensor keys.
For GRPO multi-response keys, values may be ``List[List[Tensor]]`` Stateless trait relying on ``self._data``, ``self._offsets``,
(one list of G tensors per record). These are stored as-is and ``self._num_records`` maintained by :class:`Store`.
excluded from the cumulative-length bookkeeping since they are
accessed record-by-record via ``_data`` rather than via ``fetch``.
""" """
flat_lengths = []
for key, tensors in raw.items(): def fetch_record(
self._data[key] = tensors self,
if not tensors: index: int,
self._cum[key] = [] keys: Union[str, List[str]],
flat_lengths.append(0) ):
continue return _record_fetch(self, index, keys)
# Skip nested lists (GRPO responses/masks) — record-level access
if isinstance(tensors[0], list):
self._cum[key] = [] def _record_fetch(self, index: int, keys: Union[str, List[str]]):
continue if not getattr(self, "_data", None) and self._num_records == 0:
cum = [] raise RuntimeError("Store not loaded")
total = 0 if not 0 <= index < self._num_records:
for t in tensors: raise ValueError(
total += t.shape[0] f"Record index out of bounds: {index}, num_records={self._num_records}"
cum.append(total) )
self._cum[key] = cum if isinstance(keys, str):
flat_lengths.append(cum[-1] if cum else 0) return _fetch_record_key(self, keys, index)
self._length = min(flat_lengths) if flat_lengths else 0 return {k: _fetch_record_key(self, k, index) for k in keys}
def _fetch_record_key(self, key: str, index: int):
offsets = self._offsets.get(key)
if offsets:
start = offsets[index]
end = (
offsets[index + 1]
if index + 1 < len(offsets)
else self._data[key][0].shape[0]
)
return self._data[key][0][start:end]
return self._data[key][index]
class StoreFactory(BaseFactory["Store"]): class StoreFactory(BaseFactory["Store"]):
"""Factory for creating Store instances by type name. """Factory for creating Store instances by type name."""
Example::
@StoreFactory.register("custom")
class CustomStore(Store):
...
"""
@StoreFactory.register("h5") @StoreFactory.register("h5")
class H5Store(Store): class H5Store(Store, Streamable, Recordable):
"""HDF5-based storage backend (pre-tokenized data).""" """HDF5-based storage backend (pre-tokenized data).
def load(self, path: str): 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)) self._normalize(load_h5(path))
@StoreFactory.register("bin") @StoreFactory.register("bin")
class MmapStore(Store): class MmapStore(Store, Streamable, Recordable):
"""Memory-mapped binary storage backend. """Memory-mapped binary storage backend.
Each key is a single .bin file backed by ``np.memmap(mode="r")``. Each key is a single .bin file backed by ``np.memmap(mode="r")``.
No per-process memory duplication — all DataLoader workers share the No per-process memory duplication — all DataLoader workers share the
same OS page-cache pages. same OS page-cache pages.
Format on disk:: Supports both access modes:
data_root/ - **Stream**: always available via :meth:`fetch`.
meta.json # {key: {shape, dtype}, ...} - **Record** (``fetch_record(i, key)``): only when ``meta.json``
<key>.bin # raw numpy array, one per key contains per-record ``offsets`` (written via
``save_bin(..., record_keys=...)``). Legacy bin files without
offsets have ``num_records == 0`` and ``len(store)`` reflects the
windowed sample count when ``window_size > 0``.
``segments_are_records`` is ``False`` here (bin segments are
contiguous streams, not per-record) — record access is driven
purely by ``_offsets``.
""" """
def load(self, path: str): segments_are_records = False
def __init__(
self,
window_size: int = 0,
stride: Optional[int] = None,
):
super().__init__(window_size=window_size, stride=stride)
self._mmap_refs: List[Tensor] = []
def load(self, path: str, **kwargs):
self._mmap_refs = [] self._mmap_refs = []
root = Path(path) root = Path(path)
all_raw: Dict[str, List[Tensor]] = {} all_raw: Dict[str, List[Tensor]] = {}
all_offsets: Dict[str, List[int]] = {}
meta_paths = [ meta_paths = [
Path(p) for p in glob.glob(str(root / "**" / "meta.json"), recursive=True) Path(p) for p in glob.glob(str(root / "**" / "meta.json"), recursive=True)
] ]
for meta_path in meta_paths: for meta_path in meta_paths:
raw = load_bin(str(meta_path.parent)) raw = load_bin(str(meta_path.parent))
off = load_bin_offsets(str(meta_path.parent))
for key, tensors in raw.items(): for key, tensors in raw.items():
if key not in all_raw: if key not in all_raw:
all_raw[key] = [] all_raw[key] = []
all_raw[key].extend(tensors) all_raw[key].extend(tensors)
for key, o in off.items():
if key not in all_offsets:
all_offsets[key] = []
all_offsets[key].extend(o)
if not meta_paths: if not meta_paths:
raise FileNotFoundError(f"No meta.json found under {path}") raise FileNotFoundError(f"No meta.json found under {path}")
self._normalize(all_raw) self._normalize(all_raw, offsets=all_offsets or None)
for tensors in self._data.values(): for tensors in self._data.values():
self._mmap_refs.extend(tensors) self._mmap_refs.extend(tensors)
@StoreFactory.register("jsonl") class JsonlSource:
class JsonlStore(Store): """Read raw JSON records from a ``.jsonl`` file or directory.
"""On-the-fly tokenization store for raw JSONL files.
A JSONL dataset directory contains ``*.jsonl`` files plus a A thin reader used by :class:`JsonlStore` in processor mode — holds
``dataset_config.json`` file that follows the same schema as no tokenizer, performs no tokenisation, just yields dicts.
:class:`PipelineConfig` with an additional ``tokenizer_path`` field.
Records are tokenized when the store is loaded and concatenated into
segmented tensors matching the key layout expected by the dataset
classes (``sequence``, ``loss_mask``, ``position_ids``, ...).
""" """
CONFIG_NAME = "dataset_config.json" def __init__(self, path: str):
self.path = Path(path)
self._records: Optional[List[dict]] = None
def load(self, path: str): def load(self) -> List[dict]:
root = Path(path) if self._records is None:
config_path = root / self.CONFIG_NAME self._records = self._read(self.path)
if not config_path.exists(): return self._records
raise FileNotFoundError(
f"JSONL dataset config not found: {config_path}. "
f"Expected {self.CONFIG_NAME} alongside *.jsonl files."
)
with open(config_path, "r", encoding="utf-8") as f: @staticmethod
raw_config = json.load(f) def _read(root: Path) -> List[dict]:
if root.is_file():
return JsonlSource._read_file(root)
return JsonlSource._read_dir(root)
tokenizer_path = raw_config.pop("tokenizer_path", None) or str(root) @staticmethod
self.config = PipelineConfig.from_dict(raw_config) def _read_file(path: Path) -> List[dict]:
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path) records: List[dict] = []
mask_builder = MaskBuilderFactory.create("sectioned") with open(path, "r", encoding="utf-8") as f:
position_strategy = PositionIdStrategyFactory.create(
self.config.output.position_ids_mode
)
raw: Dict[str, List[Tensor]] = {}
doc_sequences: List[List[int]] = []
def _process_item(item: dict) -> None:
nonlocal raw, doc_sequences
result = mask_builder.build(item, self.config, tokenizer)
if result is None:
return
result.pop("domain", None)
primary_ids = self._primary_ids(result)
if not primary_ids:
return
doc_sequences.append(primary_ids)
for key, ids in result.items():
if key not in raw:
raw[key] = []
if ids and isinstance(ids[0], list):
# GRPO multi-response: List[List[int]] → List[Tensor]
raw[key].append(
[torch.tensor(sub, dtype=self._infer_dtype(sub)) for sub in ids]
)
else:
raw[key].append(torch.tensor(ids, dtype=self._infer_dtype(ids)))
for jsonl_path in sorted(root.glob("*.jsonl")):
with open(jsonl_path, "r", encoding="utf-8") as f:
for line in f: for line in f:
line = line.strip() line = line.strip()
if not line: if not line:
continue continue
try: try:
item = json.loads(line) records.append(json.loads(line))
except json.JSONDecodeError: except json.JSONDecodeError:
logger.warning( logger.warning("Failed to parse JSON line in %s, skipping", path)
"Failed to parse JSON line in %s, skipping", jsonl_path return records
@staticmethod
def _read_dir(root: Path) -> List[dict]:
records: List[dict] = []
for jsonl_path in sorted(root.glob("*.jsonl")):
records.extend(JsonlSource._read_file(jsonl_path))
return records
@StoreFactory.register("jsonl")
class JsonlStore(Store, Streamable, Recordable):
"""JSONL reader with eager/lazy tokenisation modes.
A JSONL dataset is a ``.jsonl`` file or a directory of ``*.jsonl``
files plus (optionally) a ``dataset_config.json`` describing the
tokenization pipeline.
Three ways to supply an eager transform (first match wins):
- **Explicit** (``transform=``): caller-built
:class:`TokenizeTransform` applied eagerly.
- **Config file**: ``dataset_config.json`` alongside the ``*.jsonl``
files — loaded via :meth:`TokenizeTransform.from_config_file`.
- **Default messages** (``tokenizer_path=`` given, no config file):
a built-in chatml config that tokenises the ``messages`` field,
masking every role except ``assistant`` (loss on assistant only).
Lets SFT/SEQ train straight from a chat-style JSONL directory
without a hand-written config.
Two tokenisation modes, selected at :meth:`load` time:
- **Eager** (default): applies the transform to every record at load
time and registers per-key tensors via ``_normalize``. Both
``fetch`` (stream) and ``fetch_record`` (record) work.
- **Lazy** (``processor=fn`` passed): keeps raw records and defers
tokenisation to ``fetch_record``. Only record access works —
``len(store)`` returns ``num_records``; stream primitives raise.
"""
CONFIG_NAME = "dataset_config.json"
segments_are_records = True
_DEFAULT_MESSAGES_CONFIG = {
"version": 1,
"input": {
"sections": [{"field": "messages", "action": "$role", "template": True}]
},
"mask": {"system": "mask", "user": "mask", "assistant": "train"},
"mask_default": "mask",
"output": {"position_ids_mode": "doc_reset"},
}
def __init__(
self,
window_size: int = 0,
stride: Optional[int] = None,
):
super().__init__(window_size=window_size, stride=stride)
self._source: Optional[JsonlSource] = None
self._processor: Optional[Callable[[dict], Dict[str, Tensor]]] = None
self._keys_cache: Optional[List[str]] = None
def load(self, path: str, transform=None, processor=None, **kwargs):
self._source = JsonlSource(path)
records = self._source.load()
if processor is not None:
self._processor = processor
self._num_records = len(records)
return
if transform is None:
root = Path(path)
config_path = root / self.CONFIG_NAME if root.is_dir() else None
if config_path is not None and config_path.exists():
transform = TokenizeTransform.from_config_file(str(config_path))
else:
tokenizer_path = kwargs.get("tokenizer_path")
if not tokenizer_path:
raise FileNotFoundError(
f"JSONL dataset config not found. Expected "
f"{self.CONFIG_NAME} alongside *.jsonl files, pass an "
f"explicit transform, pass processor= for lazy "
f"on-the-fly tokenisation, or pass tokenizer_path= to "
f"use the built-in messages config."
) )
continue config = PipelineConfig.from_dict(self._DEFAULT_MESSAGES_CONFIG)
_process_item(item) transform = TokenizeTransform(config, tokenizer_path)
for json_path in sorted(root.glob("*.json")): transformed = transform.apply(records)
if json_path.name == self.CONFIG_NAME: self._normalize(transformed)
continue
with open(json_path, "r", encoding="utf-8") as f:
try:
data = json.load(f)
except json.JSONDecodeError:
logger.warning("Failed to parse JSON file %s, skipping", json_path)
continue
if isinstance(data, list):
for item in data:
_process_item(item)
elif isinstance(data, dict):
_process_item(data)
pos_ids = position_strategy.generate(doc_sequences) @property
if pos_ids: def keys(self) -> List[str]:
raw["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)] if self._processor is not None:
if self._keys_cache is None and self._num_records > 0:
sample = self._processor(self._source.load()[0])
self._keys_cache = list(sample.keys())
return self._keys_cache or []
return list(self._data.keys())
self._normalize(raw) def fetch_record(self, index: int, keys: Union[str, List[str]]):
if self._processor is not None:
if not 0 <= index < self._num_records:
raise ValueError(
f"Record index out of bounds: {index}, "
f"num_records={self._num_records}"
)
record = self._source.load()[index]
data = self._processor(record)
if isinstance(keys, str):
return data[keys]
return {k: data[k] for k in keys}
return _record_fetch(self, index, keys)
@staticmethod def fetch(self, begin: int, end: int, keys: Union[str, List[str]]):
def _primary_ids(result: dict) -> List[int]: if self._processor is not None:
"""Return the first flat integer list in *result* as the primary id sequence.""" raise RuntimeError(
for val in result.values(): "JsonlStore in lazy (processor) mode does not support "
if isinstance(val, list) and val and isinstance(val[0], int): "stream fetch(); use fetch_record() instead."
return val )
return [] return _stream_fetch(self, begin, end, keys)
@staticmethod def __getitem__(self, index: int) -> Dict[str, Tensor]:
def _infer_dtype(ids: List) -> torch.dtype: if self._processor is not None:
"""Infer tensor dtype from the first element of a token/value list.""" return self.fetch_record(index, self._record_keys())
if ids and isinstance(ids[0], float): return super().__getitem__(index)
return torch.float32
return torch.int32
+2 -1
View File
@@ -18,12 +18,13 @@ when available, otherwise falls back to ``torch.nn.functional.scaled_dot_product
""" """
from astrai.extension.loader import KERNEL_NAMES, is_available from astrai.extension.loader import KERNEL_NAMES, is_available
from astrai.extension.ops import attn_decode, attn_paged_decode, attn_prefill from astrai.extension.ops import attention, attn_decode, attn_paged_decode, attn_prefill
__all__ = [ __all__ = [
"attn_decode", "attn_decode",
"attn_paged_decode", "attn_paged_decode",
"attn_prefill", "attn_prefill",
"attention",
"is_available", "is_available",
"KERNEL_NAMES", "KERNEL_NAMES",
] ]
+52
View File
@@ -244,3 +244,55 @@ def attn_paged_decode(
return _torch_fallback( return _torch_fallback(
q, k, v, mask, causal_offset, scale, q_layout=li, kv_layout=1 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
+3
View File
@@ -67,7 +67,10 @@ class BaseFactory(ABC, Generic[T]):
if _get_origin(orig_base) is BaseFactory: if _get_origin(orig_base) is BaseFactory:
(arg,) = _get_args(orig_base) (arg,) = _get_args(orig_base)
cls._entries = {} cls._entries = {}
try:
cls._component_base = _resolve_type(arg, cls) cls._component_base = _resolve_type(arg, cls)
except Exception:
cls._component_base = None
return return
@classmethod @classmethod
+25 -6
View File
@@ -7,6 +7,7 @@ Subclasses may optionally consume ``token_ids`` for token-level parsing
(e.g. Harmony / VLM-style parsers). (e.g. Harmony / VLM-style parsers).
""" """
import json
import re import re
import uuid import uuid
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
@@ -117,6 +118,29 @@ def _parse_tool_call_json(json_str: str, complete: bool):
Returns ``(name, args, valid)``. Returns ``(name, args, valid)``.
""" """
if complete:
try:
obj = json.loads(json_str)
except json.JSONDecodeError:
return None, "", False
name = obj.get("name")
if not isinstance(name, str) or not name:
return None, "", False
args = obj.get("arguments")
if isinstance(args, dict):
if not args:
args = ""
else:
args = json.dumps(args, ensure_ascii=False)
args = args[1:-1].rstrip()
elif isinstance(args, list):
args = json.dumps(args, ensure_ascii=False) if args else ""
elif isinstance(args, str):
pass
else:
args = str(args) if args is not None else ""
return name, args, True
name_match = re.search(r'"name"\s*:\s*"([^"]*)"', json_str) name_match = re.search(r'"name"\s*:\s*"([^"]*)"', json_str)
if not name_match: if not name_match:
return None, "", False return None, "", False
@@ -127,8 +151,6 @@ def _parse_tool_call_json(json_str: str, complete: bool):
return name, "", True return name, "", True
raw = args_match.group(1).rstrip() raw = args_match.group(1).rstrip()
if complete and raw.endswith("}"):
raw = raw[:-1].rstrip()
if raw.startswith("{"): if raw.startswith("{"):
inner = raw[1:].rstrip() inner = raw[1:].rstrip()
if inner.endswith("}"): if inner.endswith("}"):
@@ -156,9 +178,6 @@ def _find_tool_calls(text: str, start_pos: int = 0):
break break
json_str = text[brace:end] json_str = text[brace:end]
if not _TOOL_CALL_HEAD_RE.search(json_str):
pos = end
continue
name, args, valid = _parse_tool_call_json(json_str, complete=True) name, args, valid = _parse_tool_call_json(json_str, complete=True)
if not valid or name is None: if not valid or name is None:
@@ -186,7 +205,7 @@ def _find_partial_tool_call(text: str, start_pos: int = 0):
return None return None
json_str = text[brace:] json_str = text[brace:]
if not _TOOL_CALL_HEAD_RE.search(json_str): if '"name"' not in json_str:
return None return None
name, args, valid = _parse_tool_call_json(json_str, complete=False) name, args, valid = _parse_tool_call_json(json_str, complete=False)
+40 -14
View File
@@ -300,7 +300,11 @@ class KVCache(ABC):
@abstractmethod @abstractmethod
def bind_tasks( def bind_tasks(
self, task_ids: List[str], total_len: int, device: torch.device self,
task_ids: List[str],
total_len: int,
device: torch.device,
write_positions: Optional[Tensor] = None,
) -> CacheView: ... ) -> CacheView: ...
def task_cached(self, task_id: str) -> int: def task_cached(self, task_id: str) -> int:
@@ -399,7 +403,11 @@ class PageCache(KVCache):
self._pool.record(page_table[i], prompt_ids, i) self._pool.record(page_table[i], prompt_ids, i)
def bind_tasks( def bind_tasks(
self, task_ids: List[str], total_len: int, device: torch.device self,
task_ids: List[str],
total_len: int,
device: torch.device,
write_positions: Optional[Tensor] = None,
) -> PageCacheView: ) -> PageCacheView:
page_table = self._table.table_tensor(task_ids, device) page_table = self._table.table_tensor(task_ids, device)
return PageCacheView(self._storage, page_table, total_len) return PageCacheView(self._storage, page_table, total_len)
@@ -409,28 +417,31 @@ class ContiguousCacheView(CacheView):
"""Contiguous KV-cache view for attention layers.""" """Contiguous KV-cache view for attention layers."""
def __init__( def __init__(
self, cache: "ContiguousCache", batch_indices: Tensor, total_len: int = 0 self,
cache: "ContiguousCache",
batch_indices: Tensor,
total_len: int = 0,
write_positions: Optional[Tensor] = None,
): ):
self._cache = cache self._cache = cache
self._batch_indices = batch_indices self._batch_indices = batch_indices
self._total_len = total_len self._total_len = total_len
self._write_positions = write_positions
def write(self, layer_id: int, k: Tensor, v: Tensor): def write(self, layer_id: int, k: Tensor, v: Tensor):
seq_len = k.size(1) seq_len = k.size(1)
start_pos = self._total_len - seq_len
indices = self._batch_indices 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.k[layer_id, indices, start_pos : start_pos + seq_len] = k
self._cache.v[layer_id, indices, start_pos : start_pos + seq_len] = v self._cache.v[layer_id, indices, start_pos : start_pos + seq_len] = v
new_len = start_pos + seq_len
for s in indices.tolist():
cur = self._cache._slot_len.get(s, 0)
if new_len > cur:
self._cache._slot_len[s] = new_len
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]: def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
max_len = max( max_len = self._total_len
self._cache._slot_len.get(int(s), 0) for s in self._batch_indices.tolist()
)
indices = self._batch_indices indices = self._batch_indices
k = self._cache.k[layer_id, indices, :max_len] k = self._cache.k[layer_id, indices, :max_len]
v = self._cache.v[layer_id, indices, :max_len] v = self._cache.v[layer_id, indices, :max_len]
@@ -491,9 +502,24 @@ class ContiguousCache(KVCache):
def task_extend(self, task_id: str, pos: int) -> bool: def task_extend(self, task_id: str, pos: int) -> bool:
return pos < self.max_seq_len return pos < self.max_seq_len
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)
def bind_tasks( def bind_tasks(
self, task_ids: List[str], total_len: int, device: torch.device self,
task_ids: List[str],
total_len: int,
device: torch.device,
write_positions: Optional[Tensor] = None,
) -> ContiguousCacheView: ) -> ContiguousCacheView:
slots = [self._task_slot[tid] for tid in task_ids] slots = [self._task_slot[tid] for tid in task_ids]
batch_indices = torch.tensor(slots, dtype=torch.long, device=device) batch_indices = torch.tensor(slots, dtype=torch.long, device=device)
return ContiguousCacheView(self, batch_indices, total_len) 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
)
+62 -18
View File
@@ -43,19 +43,40 @@ class Executor:
) )
task_ids = [t.task_id for t in tasks] task_ids = [t.task_id for t in tasks]
position_ids = (
torch.arange(start_pos, prompt_len, dtype=torch.long, device=self.device)
.unsqueeze(0)
.expand(batch_sz, -1)
)
input_mask = position_ids.unsqueeze(-1) >= torch.arange(
prompt_len, device=self.device
)
with torch.inference_mode(): with torch.inference_mode():
self.model( self.model(
input_ids, input_ids,
position_ids=torch.arange( input_mask=input_mask,
start_pos, prompt_len, dtype=torch.long, device=self.device position_ids=position_ids,
)
.unsqueeze(0)
.expand(batch_sz, -1),
paged_cache=self.kv_cache.bind_tasks(task_ids, prompt_len, self.device), paged_cache=self.kv_cache.bind_tasks(task_ids, prompt_len, self.device),
) )
def execute_decode(self, tasks: List[Task]) -> List[int]: def execute_decode(
self, tasks: List[Task], return_logprobs: bool = False
) -> List[int]:
"""Decode next token for each task.
Args:
return_logprobs: When ``True``, also record (and return)
the log-probability of each sampled token under the
post-strategy sampling distribution. The logprob is
appended to ``task.output_logprobs`` and the return
list becomes ``List[Tuple[int, float]]``.
Returns:
``List[int]`` of sampled token IDs, or
``List[Tuple[int, float]]`` of ``(token_id, logprob)`` when
``return_logprobs`` is ``True``.
"""
if not tasks: if not tasks:
return [] return []
@@ -68,7 +89,10 @@ class Executor:
position_ids = torch.tensor( position_ids = torch.tensor(
[t.next_pos for t in tasks], dtype=torch.long, device=self.device [t.next_pos for t in tasks], dtype=torch.long, device=self.device
) )
total_len = position_ids.max().item() + 1 total_len = max(t.next_pos for t in tasks) + 1
input_mask = position_ids[:, None, None] >= torch.arange(
total_len, device=self.device
)
task_ids = [t.task_id for t in tasks] task_ids = [t.task_id for t in tasks]
@@ -80,37 +104,57 @@ class Executor:
) )
history_lists = [] history_lists = []
mask_lists = [] history_lens = []
for t in tasks: for t in tasks:
window = t.rep_window window = t.rep_window
prompt_part = t.prompt_ids[-window:] prompt_part = t.prompt_ids[-window:]
ids = prompt_part + t.output_ids ids = prompt_part + t.output_ids
history_lists.append(ids) history_lists.append(ids)
mask_lists.append([True] * len(ids)) history_lens.append(len(ids))
max_len = max(len(h) for h in history_lists) max_len = max(history_lens) if history_lens else 0
padded_ids = torch.zeros( padded_ids = torch.zeros(
len(tasks), max_len, dtype=torch.long, device=self.device len(tasks), max_len, dtype=torch.long, device=self.device
) )
padded_mask = torch.zeros( padded_mask = torch.zeros(
len(tasks), max_len, dtype=torch.bool, device=self.device len(tasks), max_len, dtype=torch.bool, device=self.device
) )
for i, (h, m) in enumerate(zip(history_lists, mask_lists)): for i, h in enumerate(history_lists):
padded_ids[i, : len(h)] = torch.tensor( L = history_lens[i]
h, dtype=torch.long, device=self.device padded_ids[i, :L] = torch.as_tensor(h, dtype=torch.long, device=self.device)
) padded_mask[i, :L] = True
padded_mask[i, : len(m)] = torch.tensor(
m, dtype=torch.bool, device=self.device
)
with torch.inference_mode(): with torch.inference_mode():
outputs = self.model( outputs = self.model(
input_ids.unsqueeze(1), input_ids.unsqueeze(1),
paged_cache=self.kv_cache.bind_tasks(task_ids, total_len, self.device), input_mask=input_mask,
paged_cache=self.kv_cache.bind_tasks(
task_ids,
total_len,
self.device,
write_positions=position_ids,
),
position_ids=position_ids.unsqueeze(1), position_ids=position_ids.unsqueeze(1),
) )
logits = outputs["logits"][:, -1, :] logits = outputs["logits"][:, -1, :]
if return_logprobs:
tokens, logprobs = sample(
logits,
temperature=temperatures,
top_k=top_ks,
top_p=top_ps,
frequency_penalty=freq_penalties,
input_ids=padded_ids,
input_mask=padded_mask,
return_logprobs=True,
)
tokens_list = tokens.tolist()
logprobs_list = logprobs.tolist()
for t, lp in zip(tasks, logprobs_list):
t.output_logprobs.append(float(lp))
return list(zip(tokens_list, logprobs_list))
return sample( return sample(
logits, logits,
temperature=temperatures, temperature=temperatures,
+126 -17
View File
@@ -1,5 +1,6 @@
import logging import logging
import threading import threading
import uuid
from typing import Any, Dict, List, Optional, Tuple from typing import Any, Dict, List, Optional, Tuple
import torch import torch
@@ -31,26 +32,26 @@ class InferenceScheduler:
if max_seq_len is not None: if max_seq_len is not None:
self.max_seq_len = max_seq_len self.max_seq_len = max_seq_len
elif config.max_len is not None: elif config.max_position_embeddings is not None:
self.max_seq_len = config.max_len self.max_seq_len = config.max_position_embeddings
else: else:
raise ValueError( raise ValueError(
"max_seq_len must be provided either as argument " "max_seq_len must be provided either as argument "
"or in model config (config.max_len)" "or in model config (config.max_position_embeddings)"
) )
self.device = device or next(model.parameters()).device self.device = device or next(model.parameters()).device
self.dtype = dtype or next(model.parameters()).dtype self.dtype = dtype or next(model.parameters()).dtype
head_dim = config.dim // config.n_heads head_dim = config.hidden_size // config.num_attention_heads
if cache is not None: if cache is not None:
self._cache = cache self._cache = cache
else: else:
self._cache = ContiguousCache( self._cache = ContiguousCache(
config.n_layers, config.num_hidden_layers,
max_batch_size, max_batch_size,
self.max_seq_len, self.max_seq_len,
config.n_kv_heads, config.num_key_value_heads,
head_dim, head_dim,
self.device, self.device,
self.dtype, self.dtype,
@@ -138,15 +139,10 @@ class InferenceScheduler:
t.task_id, t.prompt_ids, start_logical_page t.task_id, t.prompt_ids, start_logical_page
) )
pos_groups: Dict[int, List[Task]] = {} decode_tasks = self._task_mgr.get_active_tasks()
for t in self._task_mgr.get_active_tasks():
pos_groups.setdefault(t.next_pos, []).append(t)
for next_pos in sorted(pos_groups.keys()):
group = sorted(pos_groups[next_pos], key=lambda t: t.task_id)
valid: List[Task] = [] valid: List[Task] = []
for t in group: for t in sorted(decode_tasks, key=lambda t: t.task_id):
if cache.task_extend(t.task_id, t.next_pos): if cache.task_extend(t.task_id, t.next_pos):
valid.append(t) valid.append(t)
else: else:
@@ -159,13 +155,15 @@ class InferenceScheduler:
for t, ntok in zip(valid, next_tokens): for t, ntok in zip(valid, next_tokens):
t.output_ids.append(ntok) t.output_ids.append(ntok)
t.output_tokens += 1 t.output_tokens += 1
self._task_mgr.invoke_callback( new_text = t.decode_new_token(self._task_mgr.tokenizer)
t.task_id, if new_text:
self._task_mgr.tokenizer.decode([ntok]), self._task_mgr.invoke_callback(t.task_id, new_text)
)
for t in valid: for t in valid:
if t.is_finished(stop_ids): if t.is_finished(stop_ids):
remaining = t.flush_remaining(self._task_mgr.tokenizer)
if remaining:
self._task_mgr.invoke_callback(t.task_id, remaining)
self._task_mgr.invoke_callback(t.task_id, STOP) self._task_mgr.invoke_callback(t.task_id, STOP)
except Exception as e: except Exception as e:
@@ -197,6 +195,117 @@ class InferenceScheduler:
self._cache.task_free(task.task_id) self._cache.task_free(task.task_id)
for task in self._task_mgr.get_waiting_tasks(): for task in self._task_mgr.get_waiting_tasks():
self._task_mgr.invoke_callback(task.task_id, STOP) self._task_mgr.invoke_callback(task.task_id, STOP)
self._cache.task_free(task.task_id)
self._task_mgr.clear_queues() self._task_mgr.clear_queues()
if torch.cuda.is_available(): if torch.cuda.is_available():
torch.cuda.empty_cache() torch.cuda.empty_cache()
def run_batch(
self,
prompt_ids_list: List[List[int]],
*,
max_tokens: Optional[int] = None,
temperature: float = 1.0,
top_p: float = 1.0,
top_k: int = 50,
frequency_penalty: float = 0.0,
rep_window: int = 64,
return_logprobs: bool = False,
) -> List[List[int]]:
"""Synchronous batch generation without the scheduler thread.
Accepts already-tokenized prompts (no string round-trip) and runs
prefill + decode to completion on the calling thread. Designed for
RL rollout, where logprobs of the behaviour policy must be collected
alongside generated tokens.
Args:
prompt_ids_list: ``B`` prompts, each a list of token IDs.
max_tokens: Maximum tokens to generate per prompt. ``None``
uses ``self.max_seq_len - len(prompt_ids)``.
temperature/top_p/top_k/frequency_penalty/rep_window: Sampling
parameters (uniform across the batch).
return_logprobs: If ``True``, return ``(token_ids, logprobs)``
tuples per prompt (logprobs aligned 1-to-1 with token_ids).
Returns:
``List[List[int]]`` of generated token IDs per prompt, or —
when ``return_logprobs`` is ``True`` —
``List[Tuple[List[int], List[float]]]``.
"""
stop_ids = self._task_mgr.tokenizer.stop_ids
cache = self._cache
seq_cap = self.max_seq_len
tasks: List[Task] = []
for ids in prompt_ids_list:
if len(ids) >= seq_cap:
tasks.append(None)
continue
t_max = max_tokens
if t_max is None:
t_max = seq_cap - len(ids)
else:
t_max = min(t_max, seq_cap - len(ids))
task = Task(
task_id=f"batch_{uuid.uuid4().hex[:8]}",
prompt_ids=list(ids),
max_tokens=t_max,
temperature=temperature,
top_p=top_p,
top_k=top_k,
frequency_penalty=frequency_penalty,
rep_window=rep_window,
)
if not cache.task_alloc(task.task_id, task.prompt_ids):
tasks.append(None)
continue
task.input_tokens = len(task.prompt_ids)
tasks.append(task)
try:
live = [t for t in tasks if t is not None]
prefill_groups: Dict[Tuple[int, int], List[Task]] = {}
for t in live:
key = (len(t.prompt_ids), cache.task_cached(t.task_id))
prefill_groups.setdefault(key, []).append(t)
for (prompt_len, start_pos), group in prefill_groups.items():
self._executor.execute_prefill(group, prompt_len, start_pos)
while live:
valid: List[Task] = []
for t in sorted(live, key=lambda x: x.task_id):
if cache.task_extend(t.task_id, t.next_pos):
valid.append(t)
else:
t.status = TaskStatus.ABORTED
if not valid:
break
step_out = self._executor.execute_decode(
valid, return_logprobs=return_logprobs
)
if return_logprobs:
for t, (ntok, _lp) in zip(valid, step_out):
t.output_ids.append(ntok)
t.output_tokens += 1
else:
for t, ntok in zip(valid, step_out):
t.output_ids.append(ntok)
t.output_tokens += 1
live = [t for t in valid if not t.is_finished(stop_ids)]
finally:
for t in tasks:
if t is not None:
cache.task_free(t.task_id)
results: List[Any] = []
for t in tasks:
if t is None:
results.append(([], []) if return_logprobs else [])
elif return_logprobs:
results.append((list(t.output_ids), list(t.output_logprobs)))
else:
results.append(list(t.output_ids))
return results
+63
View File
@@ -13,6 +13,40 @@ logger = logging.getLogger(__name__)
STOP = object() STOP = object()
class StreamDecoder:
"""Incremental decoder for byte-level BPE streaming.
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.
"""
__slots__ = ("_tokenizer", "_ids", "_emitted")
def __init__(self, tokenizer: AutoTokenizer):
self._tokenizer = tokenizer
self._ids: List[int] = []
self._emitted: str = ""
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 ""
class TaskStatus(Enum): class TaskStatus(Enum):
"""Task lifecycle states.""" """Task lifecycle states."""
@@ -47,10 +81,39 @@ class Task:
self.status = TaskStatus.PENDING self.status = TaskStatus.PENDING
self.output_ids: List[int] = [] self.output_ids: List[int] = []
self.output_logprobs: List[float] = []
self.input_tokens: int = 0 self.input_tokens: int = 0
self.output_tokens: int = 0 self.output_tokens: int = 0
self.arrival_time = time.time() self.arrival_time = time.time()
self.finish_time: Optional[float] = None self.finish_time: Optional[float] = None
self._decoder: Optional[StreamDecoder] = None
def decode_new_token(self, tokenizer: AutoTokenizer) -> str:
"""Decode the last appended output token, buffering incomplete
multi-byte sequences across calls.
Lazily creates a :class:`StreamDecoder` on first use.
"""
if self._decoder is None:
self._decoder = StreamDecoder(tokenizer)
return self._decoder.push(self.output_ids[-1])
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.
"""
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 @property
def next_pos(self) -> int: def next_pos(self) -> int:
+2 -2
View File
@@ -82,8 +82,8 @@ class GenerationRequest:
raise ValueError("top_k must be a non-negative integer") raise ValueError("top_k must be a non-negative integer")
if not (0.0 <= top_p <= 1.0): if not (0.0 <= top_p <= 1.0):
raise ValueError("top_p must be a float between 0.0 and 1.0") raise ValueError("top_p must be a float between 0.0 and 1.0")
if not (isinstance(temperature, (int, float)) and temperature > 0): if not (isinstance(temperature, (int, float)) and temperature >= 0):
raise ValueError("temperature must be a positive number") raise ValueError("temperature must be a non-negative number")
if not ( if not (
isinstance(frequency_penalty, (int, float)) isinstance(frequency_penalty, (int, float))
and -2.0 <= frequency_penalty <= 2.0 and -2.0 <= frequency_penalty <= 2.0
+62 -11
View File
@@ -263,6 +263,12 @@ class SamplingPipeline(BaseSamplingStrategy):
logits = strategy.apply(logits, filter_value, input_ids, input_mask) logits = strategy.apply(logits, filter_value, input_ids, input_mask)
return logits return logits
@staticmethod
def _is_greedy(temperature: Union[float, Tensor]) -> bool:
if isinstance(temperature, Tensor):
return temperature.numel() == 1 and temperature.item() == 0
return temperature == 0
@torch.inference_mode() @torch.inference_mode()
def sample( def sample(
self, self,
@@ -270,23 +276,52 @@ class SamplingPipeline(BaseSamplingStrategy):
filter_value: float = -float("inf"), filter_value: float = -float("inf"),
input_ids: Optional[Tensor] = None, input_ids: Optional[Tensor] = None,
input_mask: Optional[Tensor] = None, input_mask: Optional[Tensor] = None,
) -> Tensor: return_logprobs: bool = False,
):
"""Apply strategies then sample (softmax + multinomial). """Apply strategies then sample (softmax + multinomial).
Short-circuits to ``argmax`` when temperature is exactly 0
(deterministic / greedy decode).
Args: Args:
logits: Raw logits ``[batch, vocab_size]``. logits: Raw logits ``[batch, vocab_size]``.
input_ids: Previously generated token IDs ``[batch, seq_len]``. input_ids: Previously generated token IDs ``[batch, seq_len]``.
input_mask: Boolean mask for ``input_ids`` padding. input_mask: Boolean mask for ``input_ids`` padding.
return_logprobs: If ``True``, return ``(tokens, logprobs)``
where ``logprobs[i]`` is the log-probability of
``tokens[i]`` under the (post-strategy) sampling
distribution.
Returns: Returns:
Sampled token IDs ``[batch]``. Sampled token IDs ``[batch]``, or when ``return_logprobs``
is ``True`` a ``(token_ids, chosen_logprobs)`` tuple.
""" """
return torch.multinomial( if self._is_greedy_pipeline():
torch.softmax( tokens = logits.argmax(dim=-1)
self.apply(logits, filter_value, input_ids, input_mask), dim=-1 if not return_logprobs:
), return tokens
num_samples=1, log_probs = torch.log_softmax(logits.float(), dim=-1)
chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1)
return tokens, chosen
transformed = self.apply(logits, filter_value, input_ids, input_mask)
log_probs = torch.log_softmax(transformed.float(), dim=-1)
tokens = torch.multinomial(
torch.softmax(transformed, dim=-1), num_samples=1
).squeeze(-1) ).squeeze(-1)
if not return_logprobs:
return tokens
chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1)
return tokens, chosen
def _is_greedy_pipeline(self) -> bool:
"""True if the first strategy is greedy temperature (temp=0)."""
if not self.strategies:
return False
first = self.strategies[0]
return isinstance(first, TemperatureStrategy) and self._is_greedy(
first.temperature
)
@torch.inference_mode() @torch.inference_mode()
@@ -299,10 +334,14 @@ def sample(
input_ids: Optional[Tensor] = None, input_ids: Optional[Tensor] = None,
input_mask: Optional[Tensor] = None, input_mask: Optional[Tensor] = None,
filter_value: float = -float("inf"), filter_value: float = -float("inf"),
) -> Tensor: return_logprobs: bool = False,
):
"""Apply sampling strategies then sample (softmax + multinomial). """Apply sampling strategies then sample (softmax + multinomial).
Shortcut for ``SamplingPipeline(...).sample(logits)``. Shortcut for ``SamplingPipeline(...).sample(logits, return_logprobs=)``.
When **temperature** is exactly 0 (scalar or single-element tensor)
the function short-circuits to ``argmax`` for deterministic decode.
Args: Args:
logits: Raw logits ``[batch, vocab_size]``. logits: Raw logits ``[batch, vocab_size]``.
@@ -310,9 +349,15 @@ def sample(
(0.0 disables, range -2.0~2.0). (0.0 disables, range -2.0~2.0).
input_ids: Previously generated token IDs ``[batch, seq_len]``. input_ids: Previously generated token IDs ``[batch, seq_len]``.
input_mask: Boolean mask for ``input_ids`` padding. input_mask: Boolean mask for ``input_ids`` padding.
return_logprobs: If ``True``, also return the log-probability
of each sampled token under the (post-strategy) sampling
distribution useful for RL rollout (PPO/GRPO importance
ratios).
Returns: Returns:
Sampled token IDs ``[batch]``. Sampled token IDs ``[batch]``, or when ``return_logprobs`` is
``True`` a ``(token_ids, chosen_logprobs)`` tuple where
``chosen_logprobs`` has shape ``[batch]``.
""" """
return SamplingPipeline( return SamplingPipeline(
[ [
@@ -321,4 +366,10 @@ def sample(
TopPStrategy(top_p), TopPStrategy(top_p),
FrequencyPenaltyStrategy(frequency_penalty), FrequencyPenaltyStrategy(frequency_penalty),
] ]
).sample(logits, filter_value, input_ids, input_mask) ).sample(
logits,
filter_value=filter_value,
input_ids=input_ids,
input_mask=input_mask,
return_logprobs=return_logprobs,
)
+2 -3
View File
@@ -76,9 +76,8 @@ class GQA(nn.Module):
rotary_emb: Tensor, rotary_emb: Tensor,
attn_mask: Tensor = None, attn_mask: Tensor = None,
paged_cache: Optional[CacheView] = None, paged_cache: Optional[CacheView] = None,
is_causal: bool = False,
) -> Tensor: ) -> Tensor:
is_causal = attn_mask is None
q = self._split_heads(self.q_proj(x), self.n_heads) q = self._split_heads(self.q_proj(x), self.n_heads)
k = self._split_heads(self.k_proj(x), self.n_kv_heads) k = self._split_heads(self.k_proj(x), self.n_kv_heads)
v = self._split_heads(self.v_proj(x), self.n_kv_heads) v = self._split_heads(self.v_proj(x), self.n_kv_heads)
@@ -163,9 +162,9 @@ class MLA(nn.Module):
rotary_emb: Tensor, rotary_emb: Tensor,
attn_mask: Tensor = None, attn_mask: Tensor = None,
paged_cache: Optional[CacheView] = None, paged_cache: Optional[CacheView] = None,
is_causal: bool = False,
) -> Tensor: ) -> Tensor:
bsz, seq_len, _ = x.size() bsz, seq_len, _ = x.size()
is_causal = attn_mask is None
q = self.q_proj(x) q = self.q_proj(x)
q = q.view(bsz, seq_len, self.n_heads, self.head_dim) q = q.view(bsz, seq_len, self.n_heads, self.head_dim)
+13 -3
View File
@@ -14,10 +14,18 @@ class DecoderBlock(nn.Module):
def __init__(self, config, layer_id: int): def __init__(self, config, layer_id: int):
super().__init__() super().__init__()
cfg = asdict(config) cfg = asdict(config)
cfg["down_init_std"] = 0.02 / (2 * config.n_layers) ** 0.5 cfg.update(
dim=config.hidden_size,
dim_ffn=config.intermediate_size,
n_layers=config.num_hidden_layers,
n_heads=config.num_attention_heads,
n_kv_heads=config.num_key_value_heads,
norm_eps=config.rms_norm_eps,
down_init_std=0.02 / (2 * config.num_hidden_layers) ** 0.5,
)
self.attention = AttnFactory.create(config.attn_type, **cfg, layer_id=layer_id) self.attention = AttnFactory.create(config.attn_type, **cfg, layer_id=layer_id)
self.input_norm = RMSNorm(config.dim, config.norm_eps) self.input_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
self.post_attention_norm = RMSNorm(config.dim, config.norm_eps) self.post_attention_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
self.mlp = FFNFactory.create(config.ffn_type, **cfg) self.mlp = FFNFactory.create(config.ffn_type, **cfg)
def forward( def forward(
@@ -26,12 +34,14 @@ class DecoderBlock(nn.Module):
rotary_emb: Tensor, rotary_emb: Tensor,
attention_mask: Optional[Tensor] = None, attention_mask: Optional[Tensor] = None,
paged_cache: Optional[CacheView] = None, paged_cache: Optional[CacheView] = None,
is_causal: bool = False,
) -> Tensor: ) -> Tensor:
attn_output = self.attention( attn_output = self.attention(
self.input_norm(x), self.input_norm(x),
rotary_emb, rotary_emb,
attention_mask, attention_mask,
paged_cache, paged_cache,
is_causal,
) )
x = attn_output + x x = attn_output + x
x = self.mlp(self.post_attention_norm(x)) + x x = self.mlp(self.post_attention_norm(x)) + x
+6 -2
View File
@@ -39,8 +39,12 @@ class LoRALinear(nn.Module):
self.r = r self.r = r
self.scaling = alpha / r self.scaling = alpha / r
self.lora_A = nn.Parameter(torch.randn(r, self.weight.shape[1]) / r) device = self.weight.device
self.lora_B = nn.Parameter(torch.zeros(self.weight.shape[0], r)) dtype = self.weight.dtype
lora_a = torch.randn(r, self.weight.shape[1], device=device, dtype=dtype) / r
lora_b = torch.zeros(self.weight.shape[0], r, device=device, dtype=dtype)
self.lora_A = nn.Parameter(lora_a)
self.lora_B = nn.Parameter(lora_b)
self._merged = False self._merged = False
def forward(self, x): def forward(self, x):
+15 -7
View File
@@ -18,20 +18,28 @@ class EmbeddingEncoder(AutoModel):
def __init__(self, config: EncoderConfig): def __init__(self, config: EncoderConfig):
super().__init__(config) super().__init__(config)
self.config = config self.config = config
rope_dim = config.dim // config.n_heads rope_dim = config.hidden_size // config.num_attention_heads
rope_base = config.rope_theta if config.rope_theta is not None else 10000 rope_base = config.rope_theta if config.rope_theta is not None else 10000
self.rotary_embedding = RotaryEmbedding( self.rotary_embedding = RotaryEmbedding(
rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling rope_dim,
config.max_position_embeddings,
rope_base,
rope_scaling=config.rope_scaling,
) )
self.embed_tokens = Embedding( self.embed_tokens = Embedding(
config.vocab_size, config.dim, neftune_alpha=config.neftune_alpha config.vocab_size,
config.hidden_size,
neftune_alpha=config.neftune_alpha,
) )
self.layers = nn.ModuleList( self.layers = nn.ModuleList(
[DecoderBlock(config, layer_id) for layer_id in range(config.n_layers)] [
DecoderBlock(config, layer_id)
for layer_id in range(config.num_hidden_layers)
]
) )
self.norm = RMSNorm(config.dim, config.norm_eps) self.norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
self.pooling_type = config.pooling_type or "mean" self.pooling_type = config.pooling_type or "mean"
self.normalize_embeddings = config.normalize_embeddings or False self.normalize_embeddings = config.normalize_embeddings or False
@@ -59,10 +67,10 @@ class EmbeddingEncoder(AutoModel):
x = self.embed_tokens(input_ids) x = self.embed_tokens(input_ids)
rotary_emb = self.rotary_embedding(x, position_ids) rotary_emb = self.rotary_embedding(x, position_ids)
attn_mask = process_attention_mask(x, position_ids, input_mask, is_causal=False) attn_mask = process_attention_mask(input_mask)
for layer in self.layers: for layer in self.layers:
x = layer(x, rotary_emb, attn_mask, paged_cache=None) x = layer(x, rotary_emb, attn_mask)
hidden_states = self.norm(x) hidden_states = self.norm(x)
+26 -34
View File
@@ -15,32 +15,15 @@ from astrai.model.components.rope import RotaryEmbedding
def process_attention_mask( def process_attention_mask(
input_tensor: Tensor, input_mask: Optional[Tensor],
position_ids: Optional[Tensor],
input_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Optional[Tensor]: ) -> Optional[Tensor]:
if position_ids is None:
return None
if input_mask is not None and input_mask.dim() > 2:
return input_mask
device = input_tensor.device
B = input_tensor.size(0)
T = position_ids.max().item() + 1
if input_mask is None: if input_mask is None:
if position_ids.min().item() == 0 and is_causal:
return None return None
attend = torch.ones(B, 1, T, dtype=torch.bool, device=device) if input_mask.dim() == 2:
else: return input_mask[:, None, None, :]
attend = input_mask[:, :T].to(device=device, dtype=torch.bool).unsqueeze(1) if input_mask.dim() == 3:
return input_mask[:, None, :, :]
if is_causal: return input_mask
causal = position_ids.unsqueeze(-1) >= torch.arange(T, device=device)
attend = attend & causal
return attend.unsqueeze(1)
@AutoModel.register("autoregressive_lm") @AutoModel.register("autoregressive_lm")
@@ -53,24 +36,32 @@ class AutoRegressiveLM(AutoModel):
rope_dim = ( rope_dim = (
config.qk_rope_head_dim config.qk_rope_head_dim
if config.attn_type == "mla" if config.attn_type == "mla"
else config.dim // config.n_heads else config.hidden_size // config.num_attention_heads
) )
rope_base = config.rope_theta if config.rope_theta is not None else 10000 rope_base = config.rope_theta if config.rope_theta is not None else 10000
self.rotary_embedding = RotaryEmbedding( self.rotary_embedding = RotaryEmbedding(
rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling rope_dim,
config.max_position_embeddings,
rope_base,
rope_scaling=config.rope_scaling,
) )
self.embed_tokens = Embedding( self.embed_tokens = Embedding(
config.vocab_size, config.dim, neftune_alpha=config.neftune_alpha config.vocab_size,
config.hidden_size,
neftune_alpha=config.neftune_alpha,
) )
self.layers = nn.ModuleList( self.layers = nn.ModuleList(
[DecoderBlock(config, layer_id) for layer_id in range(config.n_layers)] [
DecoderBlock(config, layer_id)
for layer_id in range(config.num_hidden_layers)
]
) )
self.norm = RMSNorm(config.dim, config.norm_eps) self.norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
self.lm_head = Linear(config.dim, config.vocab_size) self.lm_head = Linear(config.hidden_size, config.vocab_size)
if self.config.tie_weight is True: if self.config.tie_word_embeddings is True:
self.lm_head.weight = self.embed_tokens.weight self.lm_head.weight = self.embed_tokens.weight
self.apply(self._init_weights) self.apply(self._init_weights)
@@ -85,7 +76,7 @@ class AutoRegressiveLM(AutoModel):
state_dict = dict(state_dict) state_dict = dict(state_dict)
if self.config.tie_weight is True: if self.config.tie_word_embeddings is True:
# same tensor for embed and lm_head # same tensor for embed and lm_head
if embed_key in state_dict: if embed_key in state_dict:
state_dict[lm_head_key] = state_dict[embed_key] state_dict[lm_head_key] = state_dict[embed_key]
@@ -101,7 +92,7 @@ class AutoRegressiveLM(AutoModel):
destination=destination, prefix=prefix, keep_vars=keep_vars destination=destination, prefix=prefix, keep_vars=keep_vars
) )
if self.config.tie_weight is True: if self.config.tie_word_embeddings is True:
lm_head_key = prefix + "lm_head.weight" lm_head_key = prefix + "lm_head.weight"
if lm_head_key in state_dict: if lm_head_key in state_dict:
del state_dict[lm_head_key] del state_dict[lm_head_key]
@@ -119,10 +110,11 @@ class AutoRegressiveLM(AutoModel):
x = self.embed_tokens(input_ids) x = self.embed_tokens(input_ids)
rotary_emb = self.rotary_embedding(x, position_ids) rotary_emb = self.rotary_embedding(x, position_ids)
attn_mask = process_attention_mask(x, position_ids, input_mask, is_causal=True) attn_mask = process_attention_mask(input_mask)
use_sdpa_causal_mask = attn_mask is None
for layer in self.layers: for layer in self.layers:
x = layer(x, rotary_emb, attn_mask, paged_cache) x = layer(x, rotary_emb, attn_mask, paged_cache, use_sdpa_causal_mask)
hidden_states = self.norm(x) hidden_states = self.norm(x)
logits = self.lm_head(hidden_states) logits = self.lm_head(hidden_states)
+2
View File
@@ -4,6 +4,7 @@ from astrai.parallel.executor import (
BaseExecutor, BaseExecutor,
DDPExecutor, DDPExecutor,
ExecutorFactory, ExecutorFactory,
FSDP2Executor,
FSDPExecutor, FSDPExecutor,
GradientState, GradientState,
NoneExecutor, NoneExecutor,
@@ -35,4 +36,5 @@ __all__ = [
"NoneExecutor", "NoneExecutor",
"DDPExecutor", "DDPExecutor",
"FSDPExecutor", "FSDPExecutor",
"FSDP2Executor",
] ]
+117 -19
View File
@@ -4,17 +4,22 @@ import contextlib
import logging import logging
import os import os
from contextlib import contextmanager from contextlib import contextmanager
from typing import Optional, Tuple from typing import Any, Callable, Optional, Tuple
import torch import torch
import torch.distributed as dist import torch.distributed as dist
import torch.nn as nn import torch.nn as nn
from torch.distributed.fsdp import FullStateDictConfig, StateDictType from torch.distributed.fsdp import (
FSDPModule,
FullStateDictConfig,
StateDictType,
fully_shard,
)
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.tensor import DTensor
from torch.nn.parallel import DistributedDataParallel as DDP from torch.nn.parallel import DistributedDataParallel as DDP
from torch.optim import Optimizer from torch.optim import Optimizer
from torch.optim.lr_scheduler import LRScheduler from torch.optim.lr_scheduler import LRScheduler
from torch.utils.data import DataLoader
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.parallel.setup import get_rank, get_world_size from astrai.parallel.setup import get_rank, get_world_size
@@ -86,19 +91,25 @@ class BaseExecutor:
def prepare( def prepare(
self, self,
model: nn.Module, model_fn: Callable[[], nn.Module],
optimizer: Optional[Optimizer] = None, optimizer_fn: Optional[Callable[[nn.Module], Optimizer]] = None,
dataloader: Optional[DataLoader] = None, scheduler_fn: Optional[Callable[[Optimizer], LRScheduler]] = None,
scheduler: Optional[LRScheduler] = None, before_wrap: Optional[Callable[[nn.Module], nn.Module]] = None,
) -> Tuple[ ) -> Tuple[nn.Module, Optional[Optimizer], Optional[LRScheduler]]:
nn.Module, Optional[Optimizer], Optional[DataLoader], Optional[LRScheduler] model = model_fn()
]: if before_wrap is not None:
model = before_wrap(model)
model = self._prepare_model(model) model = self._prepare_model(model)
if optimizer is not None: optimizer = None
scheduler = None
if optimizer_fn is not None:
optimizer = optimizer_fn(model)
if scheduler_fn is not None:
scheduler = scheduler_fn(optimizer)
optimizer = AccumOptimizer(optimizer, self.gradient_state) optimizer = AccumOptimizer(optimizer, self.gradient_state)
if scheduler is not None: if scheduler is not None:
scheduler = AccumScheduler(scheduler, self.gradient_state) scheduler = AccumScheduler(scheduler, self.gradient_state)
return model, optimizer, dataloader, scheduler return model, optimizer, scheduler
def _prepare_model(self, model: nn.Module) -> nn.Module: def _prepare_model(self, model: nn.Module) -> nn.Module:
return model return model
@@ -224,13 +235,6 @@ class DDPExecutor(BaseExecutor):
return model.module.state_dict() return model.module.state_dict()
return model.state_dict() return model.state_dict()
def _gather_state_dict(self, model: nn.Module):
if not self.use_distributed:
return self.unwrap_model(model)
if get_rank() != 0:
return None
return self.unwrap_model(model)
@ExecutorFactory.register("fsdp") @ExecutorFactory.register("fsdp")
class FSDPExecutor(BaseExecutor): class FSDPExecutor(BaseExecutor):
@@ -307,3 +311,97 @@ class FSDPExecutor(BaseExecutor):
return model.state_dict() return model.state_dict()
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).
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.
Original ``Parameter`` objects are preserved (as DTensors) no
``FlatParameter``, no ``use_orig_params=True`` hack.
"""
def __init__(
self,
grad_accum_steps: int = 1,
mesh: Optional[Any] = None,
mp_policy: Optional[Any] = None,
reshard_after_forward: bool = True,
):
super().__init__(grad_accum_steps=grad_accum_steps)
self._mesh = mesh
self._mp_policy = mp_policy
self._reshard_after_forward = reshard_after_forward
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")
return model
kwargs = dict(
mesh=self._mesh,
mp_policy=self._mp_policy,
reshard_after_forward=self._reshard_after_forward,
)
kwargs = {k: v for k, v in kwargs.items() if v is not None}
for child in model.children():
if isinstance(child, nn.ModuleList):
for sub in child:
fully_shard(sub, **kwargs)
else:
fully_shard(child, **kwargs)
logger.info(
"FSDP2 wrapping applied to %d direct children (root skipped for ABC compat)",
len(list(model.children())),
)
return model
@contextmanager
def _no_sync(self, model: nn.Module):
fsdp_modules = [m for m in model.modules() if isinstance(m, FSDPModule)]
if fsdp_modules:
for m in fsdp_modules:
m.set_requires_gradient_sync(False, recurse=True)
try:
yield
finally:
for m in fsdp_modules:
m.set_requires_gradient_sync(True, recurse=True)
else:
yield
def clip_grad_norm(self, model: nn.Module, max_norm: 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)
def unwrap_model(self, model: nn.Module):
if not self.use_distributed:
return model.state_dict()
if get_rank() != 0:
return None
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()
}
for module in model.modules():
if isinstance(module, FSDPModule):
module.reshard()
return result
+44 -2
View File
@@ -1,5 +1,8 @@
import logging
import os import os
import signal
import socket import socket
import threading
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from contextlib import contextmanager from contextlib import contextmanager
from functools import wraps from functools import wraps
@@ -9,6 +12,10 @@ import torch
import torch.distributed as dist import torch.distributed as dist
import torch.multiprocessing as mp import torch.multiprocessing as mp
from astrai.parallel.signal_handler import install_early_signal_handlers
logger = logging.getLogger(__name__)
def find_free_port() -> str: def find_free_port() -> str:
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
@@ -115,6 +122,7 @@ def _run_single_rank(
func: Callable, func: Callable,
kwargs: dict, kwargs: dict,
): ):
install_early_signal_handlers()
with setup_parallel( with setup_parallel(
rank=rank, rank=rank,
world_size=world_size, world_size=world_size,
@@ -155,6 +163,7 @@ class TorchrunStrategy(LaunchStrategy):
"""External orchestrator (torchrun, SLURM, K8s) — env vars pre-set.""" """External orchestrator (torchrun, SLURM, K8s) — env vars pre-set."""
def launch(self, func: Callable, **kwargs): def launch(self, func: Callable, **kwargs):
install_early_signal_handlers()
rank = int(os.environ["RANK"]) rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"]) world_size = int(os.environ["WORLD_SIZE"])
local_rank = int(os.environ.get("LOCAL_RANK", rank)) local_rank = int(os.environ.get("LOCAL_RANK", rank))
@@ -188,6 +197,7 @@ class LocalStrategy(LaunchStrategy):
_run_single_rank(0, *args) _run_single_rank(0, *args)
return return
install_early_signal_handlers()
ctx = mp.start_processes( ctx = mp.start_processes(
_run_single_rank, _run_single_rank,
args=args, args=args,
@@ -195,14 +205,46 @@ class LocalStrategy(LaunchStrategy):
start_method=self.start_method, start_method=self.start_method,
join=False, join=False,
) )
parent_stop = threading.Event()
original_handlers = {}
def _parent_handler(signum, frame):
sig = signal.Signals(signum)
logger.warning(
"Parent (pid=%d) received %s, forwarding to children...",
os.getpid(),
sig.name,
)
parent_stop.set()
for p in ctx.processes:
if p.is_alive():
p.terminate()
for sig in (signal.SIGTERM, signal.SIGINT):
prev = signal.signal(sig, _parent_handler)
if prev not in (signal.SIG_DFL, signal.SIG_IGN, None, _parent_handler):
original_handlers[sig] = prev
try: try:
while not ctx.join(): while not ctx.join() and not parent_stop.is_set():
pass pass
except BaseException: except BaseException:
logger.warning(
"Parent received unexpected exception, terminating children..."
)
for p in ctx.processes: for p in ctx.processes:
if p.is_alive():
p.terminate() p.terminate()
ctx.join()
raise raise
finally:
for sig, handler in original_handlers.items():
signal.signal(sig, handler)
for p in ctx.processes:
p.join()
ctx.join()
def _detect_launcher() -> str: def _detect_launcher() -> str:
+53
View File
@@ -0,0 +1,53 @@
import logging
import os
import signal
import threading
logger = logging.getLogger(__name__)
_early_stop = threading.Event()
_active_context = None
def _early_handler(signum: int, frame):
sig = signal.Signals(signum)
logger.warning(
"Received %s (pid=%d), requesting graceful training stop...",
sig.name,
os.getpid(),
)
_early_stop.set()
if _active_context is not None:
_active_context.request_stop()
def install_early_signal_handlers():
for sig in (signal.SIGTERM, signal.SIGINT):
signal.signal(sig, _early_handler)
_unblock_signals()
def _unblock_signals():
try:
mask = signal.pthread_sigmask(signal.SIG_BLOCK, set())
blocked = {signal.SIGTERM, signal.SIGINT} & mask
if blocked:
signal.pthread_sigmask(signal.SIG_UNBLOCK, blocked)
except (AttributeError, OSError):
pass
def register_signal_handlers(context):
global _active_context
_active_context = context
for sig in (signal.SIGTERM, signal.SIGINT):
signal.signal(sig, _early_handler)
if _early_stop.is_set():
context.request_stop()
logger.warning("Signal was received during initialization, stopping...")
def unregister_signal_handlers():
global _active_context
_active_context = None
_early_stop.clear()
+4
View File
@@ -8,12 +8,14 @@ from astrai.preprocessing.builder import (
from astrai.preprocessing.packing import ( from astrai.preprocessing.packing import (
PackingStrategy, PackingStrategy,
PackingStrategyFactory, PackingStrategyFactory,
plan_bfd,
) )
from astrai.preprocessing.pipeline import Pipeline, filter_by_length from astrai.preprocessing.pipeline import Pipeline, filter_by_length
from astrai.preprocessing.position_id import ( from astrai.preprocessing.position_id import (
PositionIdStrategy, PositionIdStrategy,
PositionIdStrategyFactory, PositionIdStrategyFactory,
) )
from astrai.preprocessing.transform import TokenizeTransform
from astrai.preprocessing.writer import ( from astrai.preprocessing.writer import (
StoreWriter, StoreWriter,
StoreWriterFactory, StoreWriterFactory,
@@ -32,5 +34,7 @@ __all__ = [
"SingleOutputMaskBuilder", "SingleOutputMaskBuilder",
"StoreWriter", "StoreWriter",
"StoreWriterFactory", "StoreWriterFactory",
"TokenizeTransform",
"filter_by_length", "filter_by_length",
"plan_bfd",
] ]
+200
View File
@@ -94,6 +94,97 @@ class SectionRenderer:
return all_ids, loss_mask return all_ids, loss_mask
def process_sections_batch(
self,
items: list[dict],
sections: list,
config,
tokenizer,
*,
is_top_level=False,
filter_text=True,
):
"""Render and tokenize a group of records with batched Rust tokenization."""
has_template = any(s.get("template") for s in sections)
is_text_config = not has_template and all(
s["action"] == "train" for s in sections
)
plans: list[list[tuple[str, str, bool]]] = []
for item in items:
plan: list[tuple[str, str, bool]] = []
first_section = True
for sec in sections:
field = sec["field"]
action = sec["action"]
use_template = sec.get("template", False)
add_special = sec.get(
"add_special_tokens", not use_template and first_section
)
if use_template:
messages = item.get(field)
if not isinstance(messages, list) or not messages:
continue
for msg in messages:
role = msg.get("role", "")
rendered = tokenizer.apply_chat_template(
[msg], tokenize=False, add_generation_prompt=False
)
plan.append(
(rendered, _resolve_action(action, role, config), False)
)
else:
text = str(item.get(field, ""))
if not text.strip():
continue
if is_text_config and filter_text:
pp = config.preprocessing
if pp.min_chars > 0 and len(text) < pp.min_chars:
continue
if len(text) > pp.max_chars:
continue
plan.append((text, action, add_special))
first_section = False
plans.append(plan)
encoded: dict[tuple[int, int], list[int]] = {}
for add_special in (False, True):
refs = [
(item_idx, unit_idx, text)
for item_idx, plan in enumerate(plans)
for unit_idx, (text, _, add) in enumerate(plan)
if add == add_special
]
if not refs:
continue
ids_batch = tokenizer.encode(
[text for _, _, text in refs], add_special_tokens=add_special
)
for (item_idx, unit_idx, _), ids in zip(refs, ids_batch):
encoded[(item_idx, unit_idx)] = ids
outputs = []
max_len = config.preprocessing.max_seq_len
for item_idx, plan in enumerate(plans):
all_ids = []
loss_mask = []
if is_top_level and has_template and tokenizer.bos_token_id is not None:
all_ids.append(tokenizer.bos_token_id)
loss_mask.append(0)
for unit_idx, (_, action, _) in enumerate(plan):
ids = encoded[(item_idx, unit_idx)]
all_ids.extend(ids)
loss_mask.extend([1 if action == "train" else 0] * len(ids))
all_ids = all_ids[:max_len]
loss_mask = loss_mask[: len(all_ids)]
if not all_ids or (is_top_level and has_template and len(all_ids) <= 1):
outputs.append((None, None))
else:
outputs.append((all_ids, loss_mask))
return outputs
def process_list_field(self, item: dict, sections: list, config, tokenizer): def process_list_field(self, item: dict, sections: list, config, tokenizer):
"""Tokenize a list-valued field, preserving per-element boundaries. """Tokenize a list-valued field, preserving per-element boundaries.
@@ -147,6 +238,42 @@ class SectionRenderer:
return None, None return None, None
return per_item_ids, per_item_masks return per_item_ids, per_item_masks
def process_list_field_batch(self, items, sections, config, tokenizer):
per_item_ids = [[] for _ in items]
per_item_masks = [[] for _ in items]
for sec in sections:
wrappers = []
owners = []
field = sec["field"]
for item_idx, item in enumerate(items):
values = item.get(field)
if not isinstance(values, list):
continue
for val in values:
if sec.get("template", False) and not isinstance(val, list):
continue
wrappers.append({field: val if isinstance(val, list) else str(val)})
owners.append(item_idx)
rendered = self.process_sections_batch(
wrappers,
[sec],
config,
tokenizer,
is_top_level=False,
filter_text=False,
)
for owner, (ids, mask) in zip(owners, rendered):
if ids:
per_item_ids[owner].append(ids)
per_item_masks[owner].append(mask)
return [
(ids, masks) if ids else (None, None)
for ids, masks in zip(per_item_ids, per_item_masks)
]
@staticmethod @staticmethod
def is_value_section(sections: list) -> bool: def is_value_section(sections: list) -> bool:
return len(sections) == 1 and sections[0].get("action") == "value" return len(sections) == 1 and sections[0].get("action") == "value"
@@ -214,6 +341,9 @@ class BaseMaskBuilder(ABC):
@abstractmethod @abstractmethod
def build(self, item: dict, config, tokenizer) -> Optional[dict]: ... def build(self, item: dict, config, tokenizer) -> Optional[dict]: ...
def build_batch(self, items: list[dict], config, tokenizer) -> list[Optional[dict]]:
return [self.build(item, config, tokenizer) for item in items]
class MaskBuilderFactory(BaseFactory["BaseMaskBuilder"]): class MaskBuilderFactory(BaseFactory["BaseMaskBuilder"]):
pass pass
@@ -248,6 +378,27 @@ class SingleOutputMaskBuilder(BaseMaskBuilder):
result["loss_mask"] = mask result["loss_mask"] = mask
return result return result
def build_batch(self, items, config, tokenizer):
sections = config.input.sections
if not sections:
return [None] * len(items)
rendered = self.renderer.process_sections_batch(
items, sections, config, tokenizer, is_top_level=True
)
results = []
for item, (ids, mask) in zip(items, rendered):
if ids is None:
results.append(None)
continue
result = {
"sequence": ids,
"domain": _extract_domain(item, config.output.domain_key),
}
if not all(m == 1 for m in mask):
result["loss_mask"] = mask
results.append(result)
return results
@MaskBuilderFactory.register("multi") @MaskBuilderFactory.register("multi")
class MultiOutputMaskBuilder(BaseMaskBuilder): class MultiOutputMaskBuilder(BaseMaskBuilder):
@@ -317,6 +468,49 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
result["domain"] = _extract_domain(item, config.output.domain_key) result["domain"] = _extract_domain(item, config.output.domain_key)
return result return result
def build_batch(self, items, config, tokenizer):
sources_spec = getattr(config.input, "sources", None)
if not sources_spec:
return [None] * len(items)
results = [{} for _ in items]
for output_key, spec in sources_spec.items():
sections = spec.get("sections", [])
if not sections:
continue
if self.renderer.is_value_section(sections):
for item, result in zip(items, results):
value = self.renderer.extract_raw_value(item, sections)
if value is not None:
result[output_key] = value
continue
mask_key = spec.get("mask_key", f"{output_key}_mask")
if spec.get("list_field", False):
rendered = self.renderer.process_list_field_batch(
items, sections, config, tokenizer
)
else:
rendered = self.renderer.process_sections_batch(
items, sections, config, tokenizer, is_top_level=True
)
for result, (ids, mask) in zip(results, rendered):
if ids is None:
continue
result[output_key] = ids
if spec.get("list_field", False) or not all(m == 1 for m in mask):
result[mask_key] = mask
elif "mask_key" in spec:
result[mask_key] = mask
return [
({**result, "domain": _extract_domain(item, config.output.domain_key)})
if result
else None
for item, result in zip(items, results)
]
@MaskBuilderFactory.register("sectioned") @MaskBuilderFactory.register("sectioned")
class SectionedMaskBuilder(BaseMaskBuilder): class SectionedMaskBuilder(BaseMaskBuilder):
@@ -335,3 +529,9 @@ class SectionedMaskBuilder(BaseMaskBuilder):
if sources_spec: if sources_spec:
return self._multi.build(item, config, tokenizer) return self._multi.build(item, config, tokenizer)
return self._single.build(item, config, tokenizer) return self._single.build(item, config, tokenizer)
def build_batch(self, items, config, tokenizer):
sources_spec = getattr(config.input, "sources", None)
if sources_spec:
return self._multi.build_batch(items, config, tokenizer)
return self._single.build_batch(items, config, tokenizer)
+124
View File
@@ -0,0 +1,124 @@
"""Shared preprocessing kernel used by both :class:`Pipeline` and
:class:`TokenizeTransform`.
The two entry points previously duplicated ~60 % of their logic:
record iteration, mask-builder invocation, primary-id extraction,
per-key accumulation, dtype inference and position-id generation.
This module factors out the common core as pure functions so that
the online (``TokenizeTransform``) and offline (``Pipeline``) paths
stay in lockstep.
"""
from itertools import chain
from typing import Dict, Iterator, List, Optional
import torch
from astrai.config.preprocess_config import PipelineConfig
from astrai.preprocessing.builder import MaskBuilderFactory
from astrai.preprocessing.position_id import PositionIdStrategyFactory
from astrai.tokenize import AutoTokenizer
def build_preprocessing_components(config: PipelineConfig, tokenizer_path: str):
"""Load tokenizer, mask builder and position-id strategy together.
Both ``Pipeline`` and ``TokenizeTransform`` need the same triple;
centralising the construction avoids drift (e.g. one path forgetting
to create the position-id strategy).
"""
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
mask_builder = MaskBuilderFactory.create("sectioned")
position_strategy = PositionIdStrategyFactory.create(
config.output.position_ids_mode
)
return tokenizer, mask_builder, position_strategy
def primary_ids(result: dict) -> List[int]:
"""Return the first flat int-list value in *result*.
Used for token counting and position-id generation when the
primary key name is not known (DPO uses ``chosen``, GRPO uses
``prompts``, SFT uses ``sequence``).
"""
for val in result.values():
if isinstance(val, list) and val and isinstance(val[0], int):
return val
return []
def infer_dtype(ids: List) -> torch.dtype:
"""Float values become float32, everything else int32."""
if ids and isinstance(ids[0], float):
return torch.float32
return torch.int32
def iter_raw_records(
records: List[dict],
mask_builder,
config: PipelineConfig,
tokenizer,
) -> Iterator[dict]:
"""Yield mask-builder output dicts for each record, skipping failures.
Drops ``domain`` from the result (callers that need it should read
it before calling this). Each yielded dict maps a key
(``sequence``, ``chosen``, ``responses``) to either a flat
``List[int]`` or a nested ``List[List[int]]`` (GRPO responses/masks).
"""
for item in records:
result = mask_builder.build(item, config, tokenizer)
if result is None:
continue
result.pop("domain", None)
if not primary_ids(result):
continue
yield result
def to_per_record_tensors(
raw: Dict[str, list],
) -> Dict[str, List[torch.Tensor]]:
"""Convert an accumulated ``{key: [per-record ids]}`` dict to tensors.
Handles three shapes transparently:
- ``List[int]`` per record (``sequence``, ``chosen``) one tensor per record.
- ``List[List[int]]`` per record (GRPO ``responses``/``masks``) one
``List[Tensor]`` per record (nested), preserving the per-response
boundary so downstream code can index responses individually.
- ``List[int]`` for the whole shard (pre-packed keys) single tensor.
The detection mirrors the previous inline logic in
``Pipeline._flush`` and ``TokenizeTransform.apply``.
"""
tensors: Dict[str, List[torch.Tensor]] = {}
for key, ids_list in raw.items():
if ids_list and isinstance(ids_list[0], list):
tensors[key] = [
[torch.tensor(sub, dtype=infer_dtype(sub)) for sub in ids]
if ids and isinstance(ids[0], list)
else torch.tensor(ids, dtype=infer_dtype(ids))
for ids in ids_list
]
else:
tensors[key] = [
torch.tensor(list(chain.from_iterable(ids_list)), dtype=torch.int32)
]
return tensors
def build_position_ids(
sequences: List[List[int]],
strategy,
) -> Optional[List[int]]:
"""Generate position ids for *sequences* using *strategy*.
Returns ``None`` when the strategy produces no ids (e.g. ``none``
mode), so callers can skip attaching the key instead of storing
an empty list.
"""
pos_ids = strategy.generate(sequences)
return pos_ids or None
+38 -30
View File
@@ -19,6 +19,43 @@ def _truncate(seq: List[int], max_len: int, mode: str) -> List[int]:
return seq[:max_len] return seq[:max_len]
def plan_bfd(
sequences: List[List[int]], max_packed_len: int, truncation_mode: str = "keep_start"
) -> List[List[int]]:
"""Best-Fit Decreasing bin packing of *sequences* into bins.
Returns a list of bins, each bin a list of original indices into
*sequences*. Bin capacities are respected on the *truncated*
length of each sequence (so a sequence longer than
*max_packed_len* counts at *max_packed_len*).
Pure index-based so callers can apply the same plan to any
aligned key (``loss_mask``, ``position_ids``).
"""
n = len(sequences)
order = sorted(range(n), key=lambda i: len(sequences[i]), reverse=True)
bins: List[List[int]] = []
bin_lengths: List[int] = []
for orig_idx in order:
seq_len = len(_truncate(sequences[orig_idx], max_packed_len, truncation_mode))
best_bin = None
best_remain = max_packed_len + 1
for i, bl in enumerate(bin_lengths):
remain = max_packed_len - bl
if seq_len <= remain < best_remain:
best_remain = remain
best_bin = i
if best_bin is not None:
bins[best_bin].append(orig_idx)
bin_lengths[best_bin] += seq_len
else:
bins.append([orig_idx])
bin_lengths.append(seq_len)
return bins
class PackingStrategy(ABC): class PackingStrategy(ABC):
"""Reorder and truncate sequences within a shard.""" """Reorder and truncate sequences within a shard."""
@@ -70,7 +107,7 @@ class BFDPacking(PackingStrategy):
sequences = keys.get("sequence", []) sequences = keys.get("sequence", [])
if not sequences: if not sequences:
return keys return keys
bins = self._plan(sequences, max_packed_len, truncation_mode) bins = plan_bfd(sequences, max_packed_len, truncation_mode)
packed: Dict[str, List[List[int]]] = {} packed: Dict[str, List[List[int]]] = {}
for k, vals in keys.items(): for k, vals in keys.items():
@@ -91,35 +128,6 @@ class BFDPacking(PackingStrategy):
result.extend(vals[i]) result.extend(vals[i])
return result return result
@staticmethod
def _plan(
sequences: List[List[int]], max_packed_len: int, truncation_mode: str
) -> List[List[int]]:
n = len(sequences)
order = sorted(range(n), key=lambda i: len(sequences[i]), reverse=True)
bins: List[List[int]] = []
bin_lengths: List[int] = []
for orig_idx in order:
seq_len = len(
_truncate(sequences[orig_idx], max_packed_len, truncation_mode)
)
best_bin = None
best_remain = max_packed_len + 1
for i, bl in enumerate(bin_lengths):
remain = max_packed_len - bl
if seq_len <= remain < best_remain:
best_remain = remain
best_bin = i
if best_bin is not None:
bins[best_bin].append(orig_idx)
bin_lengths[best_bin] += seq_len
else:
bins.append([orig_idx])
bin_lengths.append(seq_len)
return bins
@PackingStrategyFactory.register("bfd_split") @PackingStrategyFactory.register("bfd_split")
class BFDSplitPacking(BFDPacking): class BFDSplitPacking(BFDPacking):
+119 -56
View File
@@ -4,6 +4,10 @@ Composes a :class:`BaseMaskBuilder` (selected by ``input.type``) with
sharding and flush to ``.h5`` / ``.bin`` storage. Packing, position-id sharding and flush to ``.h5`` / ``.bin`` storage. Packing, position-id
generation and storage writing are each delegated to pluggable strategies, generation and storage writing are each delegated to pluggable strategies,
dispatched by configuration keys. dispatched by configuration keys.
Record iteration, mask building, primary-id extraction and per-key
accumulation are shared with :class:`TokenizeTransform` via the
:mod:`astrai.preprocessing.core` helpers.
""" """
import json import json
@@ -17,11 +21,12 @@ import torch
import tqdm import tqdm
from astrai.config.preprocess_config import PipelineConfig from astrai.config.preprocess_config import PipelineConfig
from astrai.preprocessing.builder import MaskBuilderFactory from astrai.preprocessing.core import (
build_preprocessing_components,
primary_ids,
)
from astrai.preprocessing.packing import PackingStrategyFactory from astrai.preprocessing.packing import PackingStrategyFactory
from astrai.preprocessing.position_id import PositionIdStrategyFactory
from astrai.preprocessing.writer import StoreWriterFactory from astrai.preprocessing.writer import StoreWriterFactory
from astrai.tokenize import AutoTokenizer
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -64,20 +69,21 @@ class Pipeline:
self.output_dir = output_dir self.output_dir = output_dir
self.tokenizer_path = tokenizer_path self.tokenizer_path = tokenizer_path
self.mask_builder = MaskBuilderFactory.create("sectioned") self.tokenizer, self.mask_builder, self._position_id = (
build_preprocessing_components(config, tokenizer_path)
)
self._packer = PackingStrategyFactory.create( self._packer = PackingStrategyFactory.create(
config.preprocessing.packing_strategy config.preprocessing.packing_strategy
) )
self._position_id = PositionIdStrategyFactory.create(
config.output.position_ids_mode
)
self._writer = StoreWriterFactory.create(config.output.storage_format) self._writer = StoreWriterFactory.create(config.output.storage_format)
def transform(self, item: dict) -> Optional[dict]: def transform(self, item: dict) -> Optional[dict]:
return self.mask_builder.build(item, self.config, self._tokenizer) return self.mask_builder.build(item, self.config, self.tokenizer)
def transform_batch(self, items: list[dict]) -> list[Optional[dict]]:
return self.mask_builder.build_batch(items, self.config, self.tokenizer)
def run(self): def run(self):
self._tokenizer = AutoTokenizer.from_pretrained(self.tokenizer_path)
domains: dict = defaultdict(lambda: defaultdict(list)) domains: dict = defaultdict(lambda: defaultdict(list))
total_tokens = 0 total_tokens = 0
shard_idx: dict[str, int] = defaultdict(int) shard_idx: dict[str, int] = defaultdict(int)
@@ -85,31 +91,36 @@ class Pipeline:
pp = self.config.preprocessing pp = self.config.preprocessing
for item in tqdm.tqdm( progress = tqdm.tqdm(desc="Tokenizing", unit="docs", mininterval=0.5)
self._iter_items(), desc="Tokenizing", unit="docs", mininterval=0.5 stop = False
): for items in self._iter_batches(pp.batch_size):
if pp.max_items and count >= pp.max_items: progress.update(len(items))
break
try: try:
result = self.transform(item) results = self.transform_batch(items)
except Exception: except Exception:
logger.warning( logger.warning(
"Failed to process item #%d, skipping", count + 1, exc_info=True "Failed to process batch, retrying records individually",
exc_info=True,
) )
continue results = []
for item in items:
try:
results.append(self.transform(item))
except Exception:
logger.warning(
"Failed to process item, skipping", exc_info=True
)
results.append(None)
for result in results:
if pp.max_items and count >= pp.max_items:
stop = True
break
if result is None: if result is None:
continue continue
domain = result.pop("domain", "__default__") domain = result.pop("domain", "__default__")
ids = primary_ids(result)
is_multi = bool(getattr(self.config.input, "sources", None))
if is_multi:
ids = self._primary_ids(result)
else:
ids = result.pop("sequence")
result["sequence"] = ids
if not ids: if not ids:
continue continue
@@ -125,19 +136,14 @@ class Pipeline:
self._flush(domains, shard_idx) self._flush(domains, shard_idx)
domains.clear() domains.clear()
total_tokens = 0 total_tokens = 0
if stop:
break
progress.close()
if total_tokens > 0: if total_tokens > 0:
self._flush(domains, shard_idx) self._flush(domains, shard_idx)
@staticmethod
def _primary_ids(result: dict) -> list:
"""Return the first list-valued entry in *result* as the primary id
sequence for token counting."""
for val in result.values():
if isinstance(val, list) and val and isinstance(val[0], int):
return val
return []
@staticmethod @staticmethod
def _align_bucket(bucket: dict, result: dict, ids: list): def _align_bucket(bucket: dict, result: dict, ids: list):
"""Pad previously-accumulated keys that are missing from *result*.""" """Pad previously-accumulated keys that are missing from *result*."""
@@ -162,6 +168,17 @@ class Pipeline:
continue continue
yield json.loads(line) yield json.loads(line)
def _iter_batches(self, batch_size: int):
batch_size = max(1, batch_size)
batch = []
for item in self._iter_items():
batch.append(item)
if len(batch) >= batch_size:
yield batch
batch = []
if batch:
yield batch
def _flush(self, domains, shard_idx): def _flush(self, domains, shard_idx):
for domain, keys in domains.items(): for domain, keys in domains.items():
idx = shard_idx[domain] idx = shard_idx[domain]
@@ -170,20 +187,79 @@ class Pipeline:
original_sequences = keys.get("sequence", []) original_sequences = keys.get("sequence", [])
mode = self.config.output.position_ids_mode mode = self.config.output.position_ids_mode
if mode == "doc_reset" and original_sequences: keys = self._inject_doc_reset_position_ids(keys, mode, original_sequences)
keys["position_ids"] = [list(range(len(s))) for s in original_sequences]
keys = self._packer.apply(dict(keys), pp.max_packed_len, pp.truncation_mode) keys = self._packer.apply(dict(keys), pp.max_packed_len, pp.truncation_mode)
tensors = self._to_tensors(keys)
tensors = self._inject_continuous_position_ids(
tensors, mode, keys.get("sequence", [])
)
self._writer.save(self.output_dir, domain, idx, tensors)
shard_idx[domain] = idx + 1
first_key = "sequence" if "sequence" in tensors else next(iter(tensors))
tqdm.tqdm.write(
f" saved {domain}/shard_{idx:04d} "
f"({tensors[first_key][0].numel():,} tokens)"
)
def _inject_doc_reset_position_ids(
self,
keys: Dict[str, list],
mode: str,
original_sequences: List[List[int]],
) -> Dict[str, list]:
"""Attach per-document position_ids before packing (``doc_reset``).
``doc_reset`` position ids must enter the packer so that each
packed bin concatenates the per-doc ranges in bin order. The
per-record structure ``[range(len(s)) for s in seqs]`` is required
by the packer (it concatenates per-record lists per bin); the
``PositionIdStrategy.generate`` flattens, so it cannot be used
directly here it is only consulted for the ``continuous``
post-packing path.
"""
if mode != "doc_reset" or not original_sequences:
return keys
keys["position_ids"] = [list(range(len(s))) for s in original_sequences]
return keys
def _inject_continuous_position_ids(
self,
tensors: Dict[str, List[torch.Tensor]],
mode: str,
packed_sequences: List[List[int]],
) -> Dict[str, List[torch.Tensor]]:
"""Attach a single continuous position_ids tensor after packing.
``continuous`` mode spans the whole shard (post-packing), so it
cannot participate in bin packing it is computed from the
packed sequences and appended directly to the tensor dict.
"""
if mode != "continuous" or not packed_sequences:
return tensors
pos_ids = self._position_id.generate(packed_sequences)
if pos_ids:
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
return tensors
def _to_tensors(self, keys: Dict[str, list]) -> Dict[str, List[torch.Tensor]]:
"""Convert packed per-key id lists to tensors.
Honours ``config.output.dtype`` overrides per key; falls back to
``int32``. Handles three shapes (see
:func:`astrai.preprocessing.core.to_per_record_tensors` for the
equivalent online-path helper):
- ``List[int]`` per record one tensor per record.
- ``List[List[int]]`` per record (GRPO responses/masks) one tensor
per record, inner lists flattened.
- ``List[int]`` for the whole shard (pre-packed keys) single tensor.
"""
tensors: Dict[str, List[torch.Tensor]] = {} tensors: Dict[str, List[torch.Tensor]] = {}
for key, ids_list in keys.items(): for key, ids_list in keys.items():
dt = _STR_TO_DTYPE.get( dt = _STR_TO_DTYPE.get(
self.config.output.dtype.get(key, "int32"), torch.int32 self.config.output.dtype.get(key, "int32"), torch.int32
) )
# GRPO multi-response keys store List[List[int]] per record
# (responses/masks). Rewards store List[float] per record.
# Both produce List[Tensor] (one tensor per record), but
# responses need inner flattening while rewards do not.
if ids_list and isinstance(ids_list[0], list): if ids_list and isinstance(ids_list[0], list):
tensors[key] = [ tensors[key] = [
torch.tensor( torch.tensor(
@@ -198,17 +274,4 @@ class Pipeline:
tensors[key] = [ tensors[key] = [
torch.tensor(list(chain.from_iterable(ids_list)), dtype=dt) torch.tensor(list(chain.from_iterable(ids_list)), dtype=dt)
] ]
return tensors
if mode == "continuous" and original_sequences:
pos_ids = self._position_id.generate(keys.get("sequence", []))
if pos_ids:
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
self._writer.save(self.output_dir, domain, idx, tensors)
shard_idx[domain] = idx + 1
first_key = "sequence" if "sequence" in tensors else next(iter(tensors))
tqdm.tqdm.write(
f" saved {domain}/shard_{idx:04d} "
f"({tensors[first_key][0].numel():,} tokens)"
)
+92
View File
@@ -0,0 +1,92 @@
"""Tokenization transform for JSONL record streams.
Bridges the Reader layer (``JsonlStore`` reads raw JSON records) and the
Dataset layer (expects per-record tensors). Holds the tokenizer,
mask-builder and position-id strategy together so that I/O code stays
free of model dependencies.
The record-processing core (mask building, primary-id extraction,
per-key tensorisation, position-id generation) is shared with
:class:`astrai.preprocessing.pipeline.Pipeline` via the
:mod:`astrai.preprocessing.core` helpers.
"""
import json
from pathlib import Path
from typing import Dict, List
import torch
from astrai.config.preprocess_config import PipelineConfig
from astrai.preprocessing.core import (
build_position_ids,
build_preprocessing_components,
iter_raw_records,
to_per_record_tensors,
)
class TokenizeTransform:
"""Tokenize raw JSONL record dicts into per-key tensor lists.
Owns the three preprocessing concerns that were previously inlined in
``JsonlStore``: tokenization, loss-mask construction and position-id
generation. Constructing it loads the tokenizer, so it is intentionally
cheap to pass around once built.
Args:
config: Pipeline config describing sections / masks / position mode.
tokenizer_path: Path passed to ``AutoTokenizer.from_pretrained``.
"""
def __init__(self, config: PipelineConfig, tokenizer_path: str):
self.config = config
self.tokenizer, self.mask_builder, self.position_strategy = (
build_preprocessing_components(config, tokenizer_path)
)
@classmethod
def from_config_file(cls, config_path: str) -> "TokenizeTransform":
"""Build from a ``dataset_config.json`` file path.
The config file follows :class:`PipelineConfig` schema with an
extra ``tokenizer_path`` field. When omitted, the config's
parent directory is used as the tokenizer path.
"""
root = Path(config_path).parent
with open(config_path, "r", encoding="utf-8") as f:
raw_config = json.load(f)
tokenizer_path = raw_config.pop("tokenizer_path", None) or str(root)
config = PipelineConfig.from_dict(raw_config)
return cls(config, tokenizer_path)
def apply(self, records: List[dict]) -> Dict[str, list]:
"""Tokenize a list of raw record dicts.
Returns a dict mapping key (``sequence``, ``chosen``, ``responses``,
) to a list of per-record tensors (or nested tensor lists for
multi-response keys such as GRPO ``responses``).
"""
raw: Dict[str, list] = {}
doc_sequences: List[List[int]] = []
for result in iter_raw_records(
records, self.mask_builder, self.config, self.tokenizer
):
primary = None
for val in result.values():
if isinstance(val, list) and val and isinstance(val[0], int):
primary = val
break
if primary is not None:
doc_sequences.append(primary)
for key, ids in result.items():
raw.setdefault(key, []).append(ids)
tensors = to_per_record_tensors(raw)
pos_ids = build_position_ids(doc_sequences, self.position_strategy)
if pos_ids is not None:
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
return tensors
+2
View File
@@ -19,6 +19,7 @@ from astrai.serialization.checkpoint import (
) )
from astrai.serialization.dataset import ( from astrai.serialization.dataset import (
load_bin, load_bin,
load_bin_offsets,
load_h5, load_h5,
save_bin, save_bin,
save_h5, save_h5,
@@ -37,6 +38,7 @@ __all__ = [
"save_safetensors", "save_safetensors",
"save_torch", "save_torch",
"load_bin", "load_bin",
"load_bin_offsets",
"load_h5", "load_h5",
"save_bin", "save_bin",
"save_h5", "save_h5",
+51 -4
View File
@@ -3,7 +3,7 @@
import json import json
import os import os
from pathlib import Path from pathlib import Path
from typing import Dict, List from typing import Any, Dict, List, Optional
import h5py import h5py
import numpy as np import numpy as np
@@ -50,12 +50,43 @@ def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
return tensor_group return tensor_group
def save_bin(file_path: str, tensor_group: Dict[str, List[Tensor]]): def save_bin(
file_path: str,
tensor_group: Dict[str, List[Tensor]],
record_keys: Optional[List[str]] = None,
):
"""Save tensors as memory-mapped binary files.
When *record_keys* is provided, those keys are written with per-record
cumulative offsets in ``meta.json`` so that ``MmapStore.fetch_record``
can slice individual records from the concatenated binary without
cross-record concatenation. Keys not in *record_keys* (e.g. SEQ
``sequence``) are written as a single contiguous stream without
offsets, preserving backward compatibility.
Nested keys (``List[List[Tensor]]`` such as GRPO ``responses``) are
not supported in bin format use H5 for those.
"""
os.makedirs(file_path, exist_ok=True) os.makedirs(file_path, exist_ok=True)
record_keys = set(record_keys or [])
meta = {} meta = {}
for key, tensors in tensor_group.items(): for key, tensors in tensor_group.items():
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."
)
cat = torch.cat(tensors, dim=0) cat = torch.cat(tensors, dim=0)
meta[key] = {"shape": list(cat.shape), "dtype": str(cat.dtype).split(".")[-1]} entry: Dict[str, Any] = {
"shape": list(cat.shape),
"dtype": str(cat.dtype).split(".")[-1],
}
if key in record_keys:
offsets = [0]
for t in tensors:
offsets.append(offsets[-1] + t.shape[0])
entry["offsets"] = offsets
meta[key] = entry
np.asarray(cat.cpu().numpy()).tofile(os.path.join(file_path, f"{key}.bin")) np.asarray(cat.cpu().numpy()).tofile(os.path.join(file_path, f"{key}.bin"))
with open(os.path.join(file_path, "meta.json"), "w") as f: with open(os.path.join(file_path, "meta.json"), "w") as f:
json.dump(meta, f) json.dump(meta, f)
@@ -69,8 +100,24 @@ def load_bin(file_path: str) -> Dict[str, List[Tensor]]:
arr = np.memmap( arr = np.memmap(
os.path.join(file_path, f"{key}.bin"), os.path.join(file_path, f"{key}.bin"),
dtype=info["dtype"], dtype=info["dtype"],
mode="r+", mode="c",
shape=tuple(info["shape"]), shape=tuple(info["shape"]),
) )
segments[key] = [torch.from_numpy(arr)] segments[key] = [torch.from_numpy(arr)]
return segments return segments
def load_bin_offsets(file_path: str) -> Dict[str, List[int]]:
"""Read per-record cumulative offsets from ``meta.json``.
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).
"""
with open(os.path.join(file_path, "meta.json"), "r") as f:
meta = json.load(f)
offsets: Dict[str, List[int]] = {}
for key, info in meta.items():
if "offsets" in info:
offsets[key] = info["offsets"]
return offsets
+3 -1
View File
@@ -1,8 +1,10 @@
from astrai.tokenize.chat_template import ChatTemplate, MessageType from astrai.tokenize.chat_template import ChatTemplate, MessageType
from astrai.tokenize.tokenizer import AutoTokenizer from astrai.tokenize.tokenizer import AutoTokenizer, Message, Messages
__all__ = [ __all__ = [
"AutoTokenizer", "AutoTokenizer",
"ChatTemplate", "ChatTemplate",
"MessageType", "MessageType",
"Message",
"Messages",
] ]
+14 -1
View File
@@ -1,3 +1,4 @@
from functools import cached_property
from typing import Any, Dict, List, Optional from typing import Any, Dict, List, Optional
from jinja2 import Template from jinja2 import Template
@@ -29,7 +30,19 @@ class ChatTemplate:
self.description = description self.description = description
self.default_variables = default_variables or {} self.default_variables = default_variables or {}
self.special_tokens = special_tokens or {} self.special_tokens = special_tokens or {}
self._compiled: Template = Template(template_str)
@cached_property
def _compiled(self) -> Template:
"""Lazy-compiled Jinja2 template, cached on first access.
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.
"""
return Template(self.template_str)
@classmethod @classmethod
def from_string( def from_string(
+60 -22
View File
@@ -10,6 +10,12 @@ from tokenizers import Tokenizer
from astrai.tokenize.chat_template import ChatTemplate from astrai.tokenize.chat_template import ChatTemplate
Message = Dict[str, str]
"""Single chat message with ``role`` and ``content`` keys."""
Messages = List[Message]
"""Single conversation — a list of messages."""
class AutoTokenizer: class AutoTokenizer:
"""Base tokenizer class with automatic loading support""" """Base tokenizer class with automatic loading support"""
@@ -120,7 +126,16 @@ class AutoTokenizer:
is_pretokenized: bool = False, is_pretokenized: bool = False,
add_special_tokens: bool = True, add_special_tokens: bool = True,
) -> List: ) -> List:
"""Encode text to tokens or token IDs.""" """Encode text to token IDs.
Accepts both single strings and batches:
- ``encode("hello")`` ``[123, 456]``
- ``encode(["hello", "world"])`` ``[[123, 456], [789]]``
Batches are tokenised in parallel via the Rust backend's
``encode_batch`` (uses all available CPU cores).
"""
if self._tokenizer is None: if self._tokenizer is None:
raise RuntimeError( raise RuntimeError(
"Tokenizer not initialized. Load or create a tokenizer first." "Tokenizer not initialized. Load or create a tokenizer first."
@@ -133,15 +148,13 @@ class AutoTokenizer:
add_special_tokens=add_special_tokens, add_special_tokens=add_special_tokens,
) )
return encoded.ids if out_ids else encoded.tokens return encoded.ids if out_ids else encoded.tokens
else:
encoded_list = self._tokenizer.encode_batch( encoded_list = self._tokenizer.encode_batch(
tokens, tokens,
is_pretokenized=is_pretokenized, is_pretokenized=is_pretokenized,
add_special_tokens=add_special_tokens, add_special_tokens=add_special_tokens,
) )
return [ return [encoded.ids if out_ids else encoded.tokens for encoded in encoded_list]
encoded.ids if out_ids else encoded.tokens for encoded in encoded_list
]
def decode(self, tokens: List[int], skip_special_tokens: bool = True) -> str: def decode(self, tokens: List[int], skip_special_tokens: bool = True) -> str:
"""Decode token IDs to text.""" """Decode token IDs to text."""
@@ -164,7 +177,14 @@ class AutoTokenizer:
- tokenizer.bos_token returns string - tokenizer.bos_token returns string
- tokenizer.bos_token_id returns corresponding integer ID - tokenizer.bos_token_id returns corresponding integer ID
- tokenizer.stop_ids returns list of corresponding integer IDs for all special tokens - tokenizer.stop_ids returns list of corresponding integer IDs for all special tokens
Internal/private attrs are not intercepted: during unpickle
``__dict__`` is empty, so probing ``self._special_token_map``
would recurse infinitely.
""" """
if key.startswith("_"):
raise AttributeError(key)
# Handle stop_ids - return IDs for all special tokens # Handle stop_ids - return IDs for all special tokens
if key == "stop_ids": if key == "stop_ids":
stop_ids = [] stop_ids = []
@@ -220,45 +240,63 @@ class AutoTokenizer:
def apply_chat_template( def apply_chat_template(
self, self,
messages: List[Dict[str, str]], messages: Union[Messages, List[Messages]],
system_prompt: Optional[str] = None, system_prompt: Optional[str] = None,
tokenize: bool = True, tokenize: bool = True,
add_generation_prompt: bool = True, add_generation_prompt: bool = True,
**kwargs, **kwargs,
) -> Union[str, List[int]]: ) -> Union[str, List[int], List[str], List[List[int]]]:
""" """Apply the chat template and optionally tokenize.
Apply the chat template to messages and optionally tokenize the result.
Accepts both single conversations and batches:
- ``apply_chat_template([msg1, msg2])`` ``"..."`` or ``[ids]``
- ``apply_chat_template([[msg1, msg2], [msg3]])`` ``["..", ".."]``
or ``[[ids], [ids]]``
Batches render each conversation list and tokenise all at once via
:meth:`encode` (``List[str]`` Rust parallel ``encode_batch``).
Args: Args:
messages: List of message dicts with 'role' and 'content'. messages: Single conversation (``Messages``) or batch of
system_prompt: Optional system prompt string (auto-converted to first message). conversations (``BatchMessages``).
system_prompt: Optional system prompt prepended (single mode only).
tokenize: Whether to return token IDs (True) or raw string (False). tokenize: Whether to return token IDs (True) or raw string (False).
add_generation_prompt: Whether to add the generation prompt (default: True). add_generation_prompt: Whether to add the generation prompt.
**kwargs: Additional variables to pass to the template. **kwargs: Additional template variables.
Returns: Returns:
Either the rendered string or list of token IDs. Single mode: ``str`` or ``List[int]``.
Batch mode: ``List[str]`` or ``List[List[int]]``.
Raises:
RuntimeError: If chat template is not set.
""" """
if self._chat_template is None: if self._chat_template is None:
raise RuntimeError( raise RuntimeError(
"Chat template not set. Use set_chat_template() to set a template first." "Chat template not set. Use set_chat_template() to set a template first."
) )
# Auto-convert system_prompt to first message if provided is_batch = bool(messages) and isinstance(messages[0], list)
if is_batch:
rendered = [
self._chat_template.render(
messages=msgs,
add_generation_prompt=add_generation_prompt,
**kwargs,
)
for msgs in messages
]
if tokenize:
return self.encode(rendered) # List[str] → batch encode
return rendered
# Single conversation
if system_prompt: if system_prompt:
messages = [{"role": "system", "content": system_prompt}] + list(messages) messages = [{"role": "system", "content": system_prompt}] + list(messages)
# Render the template
rendered = self._chat_template.render( rendered = self._chat_template.render(
messages=messages, messages=messages,
add_generation_prompt=add_generation_prompt, add_generation_prompt=add_generation_prompt,
**kwargs, **kwargs,
) )
if tokenize: if tokenize:
return self.encode(rendered) return self.encode(rendered)
return rendered return rendered
+421
View File
@@ -0,0 +1,421 @@
"""Online rollout runner for RL training.
Provides:
- :class:`RawRollout` generation output container (no reward yet)
- :class:`RolloutResult` a :class:`RawRollout` with rewards attached
- :class:`BaseRewardModel` pluggable reward interface
- :class:`RolloutGenerator` KV-cache-backed generation of grouped
responses + decoding (no reward); delegates the generation loop to
:class:`~astrai.inference.core.scheduler.InferenceScheduler.run_batch`
so rollout and the production inference server share one code path
- :class:`RolloutRunner` orchestrates generation + scoring with a
step-driven cache; its ``__call__`` returns ``(RolloutResult, is_fresh)``
so callers do not need to rely on object identity to detect refreshes.
"""
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Dict, List, Optional, Tuple
import torch
from torch import Tensor
from astrai.inference.core.scheduler import InferenceScheduler
@dataclass(kw_only=True)
class RawRollout:
"""Generation output before reward scoring.
Produced by :class:`RolloutGenerator`; consumed by :class:`RolloutRunner`
to assemble a :class:`RolloutResult` once rewards are attached.
Fields are designed to cover all common RL algorithms:
GRPO, PPO, Online DPO, Rejection Sampling, etc.
Fields:
prompts: Tokenized prompts, shape ``[B, P_len]``.
prompt_mask: Boolean mask for real prompt tokens, shape ``[B, P_len]``.
responses: Generated response token IDs, shape ``[B, G, R_max]``.
response_mask: Boolean mask for real (non-pad) response tokens,
shape ``[B, G, R_max]``.
logprobs_old: Per-token log-probs under the behaviour policy,
shape ``[B, G, R_max]``.
prompt_texts: Decoded prompt strings (for reward models that
need text).
response_texts: Decoded response strings, shape ``[B, G]``
(for reward models).
"""
prompts: Tensor
prompt_mask: Tensor
responses: Tensor
response_mask: Tensor
logprobs_old: Tensor
prompt_texts: List[str] = field(default_factory=list)
response_texts: List[List[str]] = field(default_factory=list)
@dataclass(kw_only=True)
class RolloutResult(RawRollout):
"""A :class:`RawRollout` with reward scoring attached.
Produced by :class:`RolloutRunner` once the :class:`BaseRewardModel`
has scored the decoded responses.
Fields:
rewards: Reward per response, shape ``[B, G]``.
"""
rewards: Tensor
class BaseRewardModel(ABC):
"""Pluggable reward model interface.
Subclasses should implement ``score()`` to return a ``[B, G]`` float
tensor of rewards. Implementations can be:
* A loaded reward model (e.g. ArmoRM, Skywork-Reward)
* An external API call
* A rule-based function (format, length, keyword matching)
"""
@abstractmethod
def score(self, prompts: List[str], responses: List[List[str]]) -> Tensor:
"""Score each generated response.
Args:
prompts: Raw prompt strings, length ``B``.
responses: Generated response strings, shape ``[B, G]``.
Returns:
Float tensor of shape ``[B, G]``.
"""
...
_PAD = 0
class RolloutGenerator:
"""Pure generation + decoding for a group of responses per prompt.
Delegates the prefill/decode loop to
:meth:`~astrai.inference.core.scheduler.InferenceScheduler.run_batch`,
which uses a real KV cache (no O() recompute). Has no dependency
on any reward model; can be reused in isolation for offline
generation, qualitative sampling, or eval pipelines.
"""
def __init__(
self,
scheduler: InferenceScheduler,
tokenizer,
max_tokens: int = 1024,
group_size: int = 8,
temperature: float = 1.0,
top_k: int = 0,
top_p: float = 1.0,
frequency_penalty: float = 0.0,
rep_window: int = 64,
):
self.scheduler = scheduler
self.tokenizer = tokenizer
self.max_tokens = max_tokens
self.group_size = group_size
self.temperature = temperature
self.top_k = top_k
self.top_p = top_p
self.frequency_penalty = frequency_penalty
self.rep_window = rep_window
@torch.no_grad()
def generate(self, batch: Dict) -> RawRollout:
"""Expand prompts by ``group_size`` and generate one response each.
Accepted batch formats (per sample, repeated B times):
- **messages**: ``{"messages": [{"role": "user", "content": "..."}, ...]}``
- **instruction + input + output**: ``{"instruction": "...",
"input": "...", "output": "..."}`` mapped to ``system`` /
``user`` / ``assistant`` messages; ``input`` and ``output``
are optional and skipped when empty.
Both are rendered through the tokenizer's chat template with
``add_generation_prompt=True`` so rollout prompts match the
format the policy was SFT-trained on.
"""
model = self.scheduler._executor.model
was_training = model.training
model.eval()
try:
return self._generate_eval(batch)
finally:
model.train(was_training)
def _generate_eval(self, batch: Dict) -> RawRollout:
prompt_texts, flat_prompt_ids = self._prepare_prompts(batch)
B = len(prompt_texts)
G = self.group_size
# Re-expand flat list to G copies per prompt for run_batch.
expanded_prompt_ids: List[List[int]] = []
for ids in flat_prompt_ids:
expanded_prompt_ids.extend([list(ids)] * G)
results = self.scheduler.run_batch(
expanded_prompt_ids,
max_tokens=self.max_tokens,
temperature=self.temperature,
top_k=self.top_k,
top_p=self.top_p,
frequency_penalty=self.frequency_penalty,
rep_window=self.rep_window,
return_logprobs=True,
)
if len(results) != B * G:
raise RuntimeError(
f"Rollout scheduler returned {len(results)} results, expected {B * G}"
)
for token_ids, logprobs in results:
if len(token_ids) != len(logprobs):
raise RuntimeError(
"Rollout scheduler returned misaligned token IDs and logprobs"
)
# Each element is (token_ids, logprobs); pad to max length.
max_len = 0
for token_ids, _lp in results:
max_len = max(max_len, len(token_ids))
max_len = max(max_len, 1)
device = self.scheduler.device
P_len = max(len(ids) for ids in flat_prompt_ids)
prompts_tensor = torch.zeros(B, P_len, dtype=torch.long, device=device)
prompt_mask = torch.zeros(B, P_len, dtype=torch.bool, device=device)
for i, ids in enumerate(flat_prompt_ids):
prompts_tensor[i, -len(ids) :] = torch.tensor(
ids, dtype=torch.long, device=device
)
prompt_mask[i, -len(ids) :] = True
responses = torch.full((B, G, max_len), _PAD, dtype=torch.long, device=device)
response_mask = torch.zeros((B, G, max_len), dtype=torch.bool, device=device)
logprobs_old = torch.zeros((B, G, max_len), dtype=torch.float, device=device)
flat_idx = 0
response_texts: List[List[str]] = [[] for _ in range(B)]
for i in range(B):
for g in range(G):
token_ids, lps = results[flat_idx]
flat_idx += 1
n = len(token_ids)
if n:
responses[i, g, :n] = torch.tensor(
token_ids, dtype=torch.long, device=device
)
response_mask[i, g, :n] = True
logprobs_old[i, g, :n] = torch.tensor(
lps, dtype=torch.float, device=device
)
response_texts[i].append(
self.tokenizer.decode(token_ids, skip_special_tokens=True)
)
return RawRollout(
prompts=prompts_tensor,
prompt_mask=prompt_mask,
responses=responses,
response_mask=response_mask,
logprobs_old=logprobs_old,
prompt_texts=prompt_texts,
response_texts=response_texts,
)
def _prepare_prompts(self, batch: Dict) -> Tuple[List[str], List[List[int]]]:
"""Render batch prompts to ``(texts, token_id_lists)``.
Returns two parallel lists of length B (number of prompts in
the batch). Dispatches by batch keys:
- ``"messages"``: treated as a pre-built message list per sample.
- ``"instruction"`` (optionally ``"input"`` and ``"output"``): mapped
to ``system`` / ``user`` / ``assistant`` messages respectively.
Both paths go through the tokenizer's chat template with
``add_generation_prompt=True``.
"""
if "messages" in batch:
messages_list = batch["messages"]
elif "instruction" in batch:
instructions = batch["instruction"]
B = len(instructions)
inputs = batch.get("input") or [""] * B
outputs = batch.get("output") or [""] * B
messages_list = [
self._instruction_to_messages(i, u, o)
for i, u, o in zip(instructions, inputs, outputs)
]
else:
raise ValueError(
"Rollout batch must contain either 'messages' or "
"'instruction' (optionally 'input'/'output'); got keys: "
f"{list(batch.keys())}"
)
try:
prompt_texts = self.tokenizer.apply_chat_template(
messages_list, tokenize=False, add_generation_prompt=True
)
if (
not isinstance(prompt_texts, list)
or len(prompt_texts) != len(messages_list)
or not all(isinstance(text, str) for text in prompt_texts)
):
raise TypeError("Tokenizer does not support batched chat templates")
flat_prompt_ids = self.tokenizer.encode(prompt_texts)
if len(flat_prompt_ids) != len(messages_list) or not all(
isinstance(ids, list) for ids in flat_prompt_ids
):
raise TypeError("Tokenizer does not support batched encoding")
except (TypeError, IndexError, KeyError):
# Keep compatibility with lightweight tokenizer adapters that only
# implement the single-conversation template API.
prompt_texts = []
flat_prompt_ids = []
for messages in messages_list:
text = self.tokenizer.apply_chat_template(
messages, tokenize=False, add_generation_prompt=True
)
ids = self.tokenizer.apply_chat_template(
messages, tokenize=True, add_generation_prompt=True
)
prompt_texts.append(text)
flat_prompt_ids.append(list(ids))
return prompt_texts, flat_prompt_ids
@staticmethod
def _instruction_to_messages(
instruction: str, inp: str = "", output: str = ""
) -> List[Dict[str, str]]:
"""Map instruction/input/output to chat messages.
Role mapping follows the convention used throughout the
preprocessing pipeline: ``instruction`` system, ``input``
user, ``output`` assistant. Empty fields are skipped so a
bare instruction produces a ``[system]`` list and the chat
template's ``add_generation_prompt`` adds the assistant header
for sampling.
"""
messages: List[Dict[str, str]] = []
if instruction:
messages.append({"role": "system", "content": instruction})
if inp:
messages.append({"role": "user", "content": inp})
if output:
messages.append({"role": "assistant", "content": output})
return messages
class RolloutRunner:
"""Produces :class:`RolloutResult` from a prompt batch.
Composes a :class:`RolloutGenerator` (generation + decoding) with a
:class:`BaseRewardModel` (scoring). Maintains an internal cache so
the same batch prompt can be replayed for multiple gradient steps.
A new rollout is triggered every ``rollout_interval`` calls to
:meth:`step` (or after :meth:`clear_cache`).
The ``__call__`` contract returns a ``(RolloutResult, is_fresh)``
tuple callers must use the boolean to detect a refreshed rollout
rather than relying on object identity.
Usage::
generator = RolloutGenerator(policy, tokenizer, pipeline, ...)
runner = RolloutRunner(generator, reward_model, rollout_interval=512)
result, is_fresh = runner(prompt_batch)
if is_fresh:
... # e.g. sync behaviour policy
"""
def __init__(
self,
generator: RolloutGenerator,
reward_model: BaseRewardModel,
rollout_interval: int = 512,
):
self.generator = generator
self.reward_model = reward_model
self.rollout_interval = rollout_interval
self._cache: Optional[RolloutResult] = None
self._cache_key = None
self._steps_since_rollout: int = 0
def step(self):
"""Advance the internal counter (call once per optimizer step)."""
self._steps_since_rollout += 1
def clear_cache(self):
"""Force next call to re-run rollout."""
self._cache = None
self._cache_key = None
@staticmethod
def _batch_key(batch: Dict):
"""Build a stable key for the prompt fields accepted by the generator."""
def freeze(value):
if isinstance(value, dict):
return tuple(sorted((key, freeze(val)) for key, val in value.items()))
if isinstance(value, (list, tuple)):
return tuple(freeze(item) for item in value)
return value
fields = ("messages", "instruction", "input", "output")
return tuple(
(field, freeze(batch[field])) for field in fields if field in batch
)
def _score(self, raw: RawRollout) -> RolloutResult:
rewards = self.reward_model.score(raw.prompt_texts, raw.response_texts)
if not isinstance(rewards, Tensor):
rewards = torch.as_tensor(rewards, dtype=torch.float32)
expected_shape = raw.responses.shape[:2]
if rewards.shape != expected_shape:
raise ValueError(
f"Reward model returned shape {tuple(rewards.shape)}, "
f"expected {tuple(expected_shape)}"
)
if not torch.isfinite(rewards).all():
raise ValueError("Reward model returned non-finite values")
device = raw.prompts.device
return RolloutResult(
prompts=raw.prompts,
prompt_mask=raw.prompt_mask,
responses=raw.responses,
response_mask=raw.response_mask,
rewards=rewards.to(device=device),
logprobs_old=raw.logprobs_old,
prompt_texts=raw.prompt_texts,
response_texts=raw.response_texts,
)
def __call__(self, batch: Dict[str, Tensor]) -> Tuple[RolloutResult, bool]:
"""Return ``(cached or fresh) RolloutResult`` plus an ``is_fresh`` flag.
Triggers a new rollout when ``_steps_since_rollout >= rollout_interval``
or when the cache is empty.
"""
cache_key = self._batch_key(batch)
if (
self._cache is None
or cache_key != self._cache_key
or self._steps_since_rollout >= self.rollout_interval
):
raw = self.generator.generate(batch)
self._cache = self._score(raw)
self._cache_key = cache_key
self._steps_since_rollout = 0
return self._cache, True
return self._cache, False
+159 -19
View File
@@ -9,6 +9,7 @@ import torch.nn.functional as F
from torch import Tensor from torch import Tensor
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.trainer.rollout import RolloutResult
def create_ref_model( def create_ref_model(
@@ -28,9 +29,10 @@ def move_to_device(batch: Dict[str, Tensor], device: str) -> Dict[str, Tensor]:
def get_logprobs( def get_logprobs(
model: Union[nn.Module, Callable[..., Dict[str, Tensor]]], model: nn.Module,
input_ids: Tensor, input_ids: Tensor,
mask: Tensor, attn_mask: Tensor,
loss_mask: Tensor,
reduction: str, reduction: str,
) -> Tensor: ) -> Tensor:
"""Compute token-wise log probabilities from model outputs. """Compute token-wise log probabilities from model outputs.
@@ -38,7 +40,8 @@ def get_logprobs(
Args: Args:
model: The language model model: The language model
input_ids: Input token IDs of shape [batch_size, seq_len] input_ids: Input token IDs of shape [batch_size, seq_len]
mask: Attention mask of shape [batch_size, seq_len] attn_mask: Attention mask passed to the model (may include causal).
loss_mask: Per-token mask for loss reduction.
reduction: How to reduce over sequence dimension ("mean", "sum", "none") reduction: How to reduce over sequence dimension ("mean", "sum", "none")
Returns: Returns:
@@ -51,9 +54,12 @@ def get_logprobs(
) )
shifted_input_ids = input_ids[:, 1:] shifted_input_ids = input_ids[:, 1:]
shifted_mask = mask[:, 1:] shifted_loss_mask = loss_mask[:, 1:]
logits = model(input_ids[:, :-1], mask[:, :-1])["logits"] logits = model(
input_ids[:, :-1],
attn_mask[:, :, :-1, :-1] if attn_mask.dim() == 4 else attn_mask[:, :-1],
)["logits"]
log_probs = torch.log_softmax(logits.float(), dim=-1) log_probs = torch.log_softmax(logits.float(), dim=-1)
token_logprobs = torch.gather( token_logprobs = torch.gather(
@@ -61,13 +67,13 @@ def get_logprobs(
).squeeze(-1) ).squeeze(-1)
if reduction == "mean": if reduction == "mean":
return (token_logprobs * shifted_mask).sum(dim=-1) / shifted_mask.sum( return (token_logprobs * shifted_loss_mask).sum(dim=-1) / shifted_loss_mask.sum(
dim=-1 dim=-1
).clamp(min=1.0) ).clamp(min=1.0)
elif reduction == "sum": elif reduction == "sum":
return (token_logprobs * shifted_mask).sum(dim=-1) return (token_logprobs * shifted_loss_mask).sum(dim=-1)
else: else:
return token_logprobs * shifted_mask return token_logprobs * shifted_loss_mask
def make_doc_boundary_mask(position_ids: Tensor) -> Tensor: def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
@@ -87,7 +93,15 @@ def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
class BaseStrategy(ABC): class BaseStrategy(ABC):
"""Abstract base class for training strategies.""" """Abstract base class for training strategies.
When a :class:`~astrai.trainer.rollout.RolloutRunner` is injected via
:meth:`set_rollout_runner`, the strategy transparently switches to
online mode: each ``__call__`` produces a :class:`RolloutResult`,
converts it to a training batch via :meth:`prepare_from_rollout`, and
then computes the loss. Without a runner the strategy runs in
offline mode and consumes the batch directly.
"""
def __init__( def __init__(
self, self,
@@ -98,8 +112,8 @@ class BaseStrategy(ABC):
self.model = model self.model = model
self.device = device self.device = device
self.executor = kwargs.pop("executor", None) self.executor = kwargs.pop("executor", None)
self.model_fn = kwargs.pop("model_fn", None)
self.extra_kwargs = kwargs self.extra_kwargs = kwargs
self._rollout_runner = None
@abstractmethod @abstractmethod
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor: def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
@@ -113,10 +127,54 @@ class BaseStrategy(ABC):
""" """
raise NotImplementedError raise NotImplementedError
def supports_online(self) -> bool:
"""Whether this strategy can operate with a rollout runner.
Base implementation returns ``False``; strategies that implement
:meth:`prepare_from_rollout` should override to return ``True``.
"""
return False
def set_rollout_runner(self, runner):
"""Inject a :class:`RolloutRunner` to enable online rollout mode."""
self._rollout_runner = runner
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
"""Map a :class:`RolloutResult` to the batch layout expected by
:meth:`compute_loss`.
Strategies that return ``True`` from :meth:`supports_online` must
override this. Default raises :class:`NotImplementedError`.
"""
raise NotImplementedError(
f"{type(self).__name__} does not support online rollout"
)
def _on_rollout_refresh(self):
"""Hook fired when a fresh rollout result is produced.
Override to refresh stale state (e.g. syncing the behaviour
policy). Default is a no-op.
"""
pass
def on_optimizer_step(self):
"""Advance online rollout state after a successful optimizer step."""
if self._rollout_runner is not None:
self._rollout_runner.step()
def __call__(self, batch: Dict[str, Tensor]) -> Tensor: def __call__(self, batch: Dict[str, Tensor]) -> Tensor:
"""Allow calling strategy directly as a callable.""" """Run offline or online forward depending on runner injection."""
if self._rollout_runner is None:
return self.compute_loss(batch) return self.compute_loss(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)
class StrategyFactory(BaseFactory["BaseStrategy"]): class StrategyFactory(BaseFactory["BaseStrategy"]):
"""Factory class for creating training strategy instances. """Factory class for creating training strategy instances.
@@ -225,7 +283,7 @@ class DPOStrategy(BaseStrategy):
device: str, device: str,
ref_model: nn.Module, ref_model: nn.Module,
beta: float = 0.1, beta: float = 0.1,
reduction: str = "mean", reduction: str = "sum",
**kwargs, **kwargs,
): ):
super().__init__(model, device, **kwargs) super().__init__(model, device, **kwargs)
@@ -239,13 +297,31 @@ class DPOStrategy(BaseStrategy):
chosen_mask, rejected_mask = batch["chosen_mask"], batch["rejected_mask"] chosen_mask, rejected_mask = batch["chosen_mask"], batch["rejected_mask"]
concat_ids = torch.cat([chosen_ids, rejected_ids], dim=0) concat_ids = torch.cat([chosen_ids, rejected_ids], dim=0)
concat_mask = torch.cat([chosen_mask, rejected_mask], dim=0) concat_loss_mask = torch.cat([chosen_mask, rejected_mask], dim=0)
log_pi = get_logprobs(self.model, concat_ids, concat_mask, self.reduction) # Build full attention mask: key-padding + causal
key_pad = concat_ids.bool()[:, None, None, :] # [B*2, 1, 1, S]
S = key_pad.shape[-1]
causal = torch.tril(
torch.ones(S, S, dtype=torch.bool, device=concat_ids.device)
)[None, None, :, :] # [1, 1, S, S]
full_mask = key_pad & causal # [B*2, 1, S, S] — composed
log_pi = get_logprobs(
self.model,
concat_ids,
full_mask,
concat_loss_mask,
self.reduction,
)
with torch.no_grad(): with torch.no_grad():
log_ref = get_logprobs( log_ref = get_logprobs(
self.ref_model, concat_ids, concat_mask, self.reduction self.ref_model,
concat_ids,
full_mask,
concat_loss_mask,
self.reduction,
) )
log_pi_chosen = log_pi[: chosen_ids.shape[0]] log_pi_chosen = log_pi[: chosen_ids.shape[0]]
@@ -261,6 +337,29 @@ class DPOStrategy(BaseStrategy):
return dpo_loss return dpo_loss
def supports_online(self) -> bool:
return True
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
"""Pick best/worst response per prompt by reward as chosen/rejected."""
rewards = result.rewards
responses = result.responses
masks = result.response_mask
best = rewards.argmax(dim=-1)
worst = rewards.argmin(dim=-1)
B = responses.shape[0]
idx = torch.arange(B, device=responses.device)
chosen = responses[idx, best]
chosen_mask = masks[idx, best].float()
rejected = responses[idx, worst]
rejected_mask = masks[idx, worst].float()
return {
"chosen": chosen,
"chosen_mask": chosen_mask,
"rejected": rejected,
"rejected_mask": rejected_mask,
}
@StrategyFactory.register("grpo") @StrategyFactory.register("grpo")
class GRPOStrategy(BaseStrategy): class GRPOStrategy(BaseStrategy):
@@ -315,6 +414,12 @@ class GRPOStrategy(BaseStrategy):
responses_flat = responses.view(-1, response_len) responses_flat = responses.view(-1, response_len)
masks_flat = masks.view(-1, response_len) masks_flat = masks.view(-1, response_len)
prompt_expanded = prompts.unsqueeze(1).repeat(1, group_size, 1).flatten(0, 1) prompt_expanded = prompts.unsqueeze(1).repeat(1, group_size, 1).flatten(0, 1)
prompt_mask = batch.get("prompt_mask")
if prompt_mask is None:
prompt_mask = prompts.ne(0)
prompt_mask_expanded = (
prompt_mask.unsqueeze(1).expand(-1, group_size, -1).flatten(0, 1)
)
prompt_len = prompt_expanded.size(1) prompt_len = prompt_expanded.size(1)
full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-1) full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-1)
@@ -322,20 +427,32 @@ class GRPOStrategy(BaseStrategy):
# response tokens. get_logprobs shifts the mask by one position, so # response tokens. get_logprobs shifts the mask by one position, so
# the first response token's logprob (predicted from the last prompt # the first response token's logprob (predicted from the last prompt
# token) is correctly included. # token) is correctly included.
full_masks = torch.cat([torch.zeros_like(prompt_expanded), masks_flat], dim=-1) full_masks = torch.cat(
[torch.zeros_like(prompt_expanded, dtype=torch.bool), masks_flat], dim=-1
)
# Build full attention mask: key-padding + causal
key_pad = torch.cat([prompt_mask_expanded, masks_flat.bool()], dim=-1)[
:, None, None, :
]
S = key_pad.shape[-1]
causal = torch.tril(
torch.ones(S, S, dtype=torch.bool, device=full_sequences.device)
)[None, None, :, :]
attn_mask = key_pad & causal
# get_logprobs returns [B*G, S-1] (S = prompt_len + response_len). # get_logprobs returns [B*G, S-1] (S = prompt_len + response_len).
# Response token logprobs occupy the last ``response_len`` positions # Response token logprobs occupy the last ``response_len`` positions
# (the first response token is predicted from the last prompt token). # (the first response token is predicted from the last prompt token).
token_log_probs_policy = get_logprobs( token_log_probs_policy = get_logprobs(
self.model, full_sequences, full_masks, "none" self.model, full_sequences, attn_mask, full_masks, "none"
)[:, prompt_len - 1 :] )[:, prompt_len - 1 :]
with torch.no_grad(): with torch.no_grad():
token_log_probs_old = get_logprobs( token_log_probs_old = get_logprobs(
self.old_model, full_sequences, full_masks, "none" self.old_model, full_sequences, attn_mask, full_masks, "none"
)[:, prompt_len - 1 :] )[:, prompt_len - 1 :]
token_log_probs_ref = get_logprobs( token_log_probs_ref = get_logprobs(
self.ref_model, full_sequences, full_masks, "none" self.ref_model, full_sequences, attn_mask, full_masks, "none"
)[:, prompt_len - 1 :] )[:, prompt_len - 1 :]
# Reshape to [B, G, response_len] # Reshape to [B, G, response_len]
@@ -372,3 +489,26 @@ class GRPOStrategy(BaseStrategy):
total_loss = policy_loss + kl_penalty total_loss = policy_loss + kl_penalty
return total_loss return total_loss
def supports_online(self) -> bool:
return True
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
return {
"prompts": result.prompts,
"prompt_mask": result.prompt_mask,
"responses": result.responses,
"masks": result.response_mask,
"rewards": result.rewards,
}
def _on_rollout_refresh(self):
"""Sync the behaviour policy whenever a fresh rollout arrives."""
self.sync_old_model()
# Factory aliases: online variants use the same strategy class; the
# ``RolloutRunner`` is injected by ``TrainContextBuilder`` to enable
# online mode, so no separate subclass is needed.
StrategyFactory._entries["online_grpo"] = GRPOStrategy
StrategyFactory._entries["online_dpo"] = DPOStrategy
+118 -55
View File
@@ -1,3 +1,4 @@
import threading
from dataclasses import dataclass, field from dataclasses import dataclass, field
from pathlib import Path from pathlib import Path
from typing import Any, Dict, Optional, Self from typing import Any, Dict, Optional, Self
@@ -7,12 +8,15 @@ import torch.nn as nn
from torch.utils.data import DataLoader, random_split from torch.utils.data import DataLoader, random_split
from astrai.config.train_config import TrainConfig from astrai.config.train_config import TrainConfig
from astrai.dataset import ResumableDistributedSampler from astrai.dataset import RDSampler
from astrai.inference.core.scheduler import InferenceScheduler
from astrai.model.components.lora import inject_lora from astrai.model.components.lora import inject_lora
from astrai.parallel.executor import BaseExecutor, ExecutorFactory from astrai.parallel.executor import BaseExecutor, ExecutorFactory
from astrai.parallel.setup import get_current_device, get_rank, get_world_size from astrai.parallel.setup import get_current_device, get_rank, get_world_size
from astrai.protocols import OptimizerProtocol, SchedulerProtocol from astrai.protocols import OptimizerProtocol, SchedulerProtocol
from astrai.serialization import Checkpoint, load_json from astrai.serialization import Checkpoint, load_json
from astrai.tokenize import AutoTokenizer
from astrai.trainer.rollout import RolloutGenerator, RolloutRunner
from astrai.trainer.strategy import BaseStrategy, StrategyFactory, create_ref_model from astrai.trainer.strategy import BaseStrategy, StrategyFactory, create_ref_model
@@ -27,7 +31,6 @@ class TrainContext:
config: TrainConfig = field(default=None) config: TrainConfig = field(default=None)
model_config: dict = field(default_factory=dict) model_config: dict = field(default_factory=dict)
executor: BaseExecutor = field(default=None) executor: BaseExecutor = field(default=None)
epoch: int = field(default=0) epoch: int = field(default=0)
consumed_samples: int = field(default=0) consumed_samples: int = field(default=0)
loss: float = field(default=0.0) loss: float = field(default=0.0)
@@ -39,6 +42,15 @@ class TrainContext:
rank: int = field(default=0) rank: int = field(default=0)
kwargs: Dict[str, Any] = field(default_factory=dict) kwargs: Dict[str, Any] = field(default_factory=dict)
_stop_event: threading.Event = field(default_factory=threading.Event)
@property
def stop_requested(self) -> bool:
return self._stop_event.is_set()
def request_stop(self) -> None:
self._stop_event.set()
@property @property
def optimizer_step(self) -> int: def optimizer_step(self) -> int:
return self.consumed_samples // ( return self.consumed_samples // (
@@ -72,61 +84,70 @@ class TrainContextBuilder:
**cfg.executor_kwargs, **cfg.executor_kwargs,
) )
model = cfg.model_fn()
model = model.to(device=device)
model_config = {} model_config = {}
if self._param_path: if self._param_path:
config_path = Path(self._param_path) / "config.json" config_path = Path(self._param_path) / "config.json"
if config_path.exists(): if config_path.exists():
model_config = load_json(config_path) model_config = load_json(config_path)
if not model_config and hasattr(model, "config"): preloaded_state_dict = None
model_config = model.config.to_dict() preloaded_epoch = cfg.start_epoch
preloaded_consumed = cfg.start_samples * get_world_size()
preloaded_checkpoint = None
if self._param_path:
checkpoint = Checkpoint.load_any(self._param_path)
if checkpoint is not None:
preloaded_state_dict = checkpoint.state_dict
if checkpoint.config:
model_config = checkpoint.config
if self._resume:
preloaded_epoch = checkpoint.epoch 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_checkpoint = checkpoint
if not model_config and hasattr(cfg.model_fn(), "config"):
model_config = cfg.model_fn().config.to_dict()
def _before_wrap(m):
m = m.to(device=device)
if cfg.lora is not None:
inject_lora(
m,
r=cfg.lora.r,
alpha=cfg.lora.alpha,
target_modules=set(cfg.lora.target_modules),
)
if preloaded_state_dict is not None:
m.load_state_dict(preloaded_state_dict, strict=False)
return m
context = TrainContext( context = TrainContext(
model=model,
world_size=get_world_size(), world_size=get_world_size(),
rank=get_rank(), rank=get_rank(),
config=cfg, config=cfg,
model_config=model_config, model_config=model_config,
executor=executor, executor=executor,
epoch=preloaded_epoch,
consumed_samples=preloaded_consumed,
checkpoint=preloaded_checkpoint,
) )
if self._param_path: context.model, context.optimizer, context.scheduler = executor.prepare(
checkpoint = Checkpoint.load_any(self._param_path) cfg.model_fn,
if checkpoint is not None: cfg.optimizer_fn,
model.load_state_dict(checkpoint.state_dict, strict=False) cfg.scheduler_fn,
if checkpoint.config: before_wrap=_before_wrap,
context.model_config = checkpoint.config
if self._resume:
context.epoch = checkpoint.epoch or cfg.start_epoch
if checkpoint.consumed_samples > 0:
per_step = (
cfg.batch_per_device
* context.world_size
* cfg.grad_accum_steps
) )
context.consumed_samples = (
checkpoint.consumed_samples // per_step
) * per_step
else:
context.consumed_samples = (
cfg.start_samples * context.world_size
)
context.checkpoint = checkpoint
if cfg.lora is not None:
inject_lora(
model,
r=cfg.lora.r,
alpha=cfg.lora.alpha,
target_modules=set(cfg.lora.target_modules),
)
context.optimizer = cfg.optimizer_fn(model)
context.scheduler = cfg.scheduler_fn(context.optimizer)
train_dataset = cfg.dataset train_dataset = cfg.dataset
val_dataset = cfg.val_dataset val_dataset = cfg.val_dataset
@@ -141,7 +162,7 @@ class TrainContextBuilder:
) )
sampler_offset = context.consumed_samples // context.world_size sampler_offset = context.consumed_samples // context.world_size
sampler = ResumableDistributedSampler( sampler = RDSampler(
data_source=train_dataset, data_source=train_dataset,
start_epoch=context.epoch, start_epoch=context.epoch,
start_iter=sampler_offset, start_iter=sampler_offset,
@@ -154,10 +175,11 @@ class TrainContextBuilder:
num_workers=cfg.num_workers, num_workers=cfg.num_workers,
pin_memory=cfg.pin_memory, pin_memory=cfg.pin_memory,
prefetch_factor=cfg.prefetch_factor, prefetch_factor=cfg.prefetch_factor,
collate_fn=cfg.collate_fn,
) )
if val_dataset is not None: if val_dataset is not None:
val_sampler = ResumableDistributedSampler( val_sampler = RDSampler(
data_source=val_dataset, data_source=val_dataset,
start_epoch=0, start_epoch=0,
start_iter=0, start_iter=0,
@@ -171,15 +193,7 @@ class TrainContextBuilder:
num_workers=cfg.num_workers, num_workers=cfg.num_workers,
pin_memory=cfg.pin_memory, pin_memory=cfg.pin_memory,
prefetch_factor=cfg.prefetch_factor, prefetch_factor=cfg.prefetch_factor,
) collate_fn=cfg.collate_fn,
context.model, context.optimizer, context.dataloader, context.scheduler = (
executor.prepare(
model,
context.optimizer,
context.dataloader,
context.scheduler,
)
) )
if context.checkpoint and context.checkpoint.extra: if context.checkpoint and context.checkpoint.extra:
@@ -192,13 +206,22 @@ class TrainContextBuilder:
strategy_kwargs = dict(cfg.extra_kwargs) strategy_kwargs = dict(cfg.extra_kwargs)
if cfg.strategy in ("dpo", "grpo"): needs_ref = cfg.strategy in (
"dpo",
"grpo",
"online_grpo",
"online_dpo",
)
needs_old = cfg.strategy in ("grpo", "online_grpo")
if needs_ref:
ref_model = create_ref_model( ref_model = create_ref_model(
cfg.model_fn, executor.unwrap_model(context.model) cfg.model_fn, executor.unwrap_model(context.model)
).to(device=device) ).to(device=device)
strategy_kwargs["ref_model"] = ref_model strategy_kwargs["ref_model"] = ref_model
if cfg.strategy == "grpo": old_model = None
if needs_old:
old_model = create_ref_model( old_model = create_ref_model(
cfg.model_fn, executor.unwrap_model(context.model) cfg.model_fn, executor.unwrap_model(context.model)
).to(device=device) ).to(device=device)
@@ -209,8 +232,48 @@ class TrainContextBuilder:
model=context.model, model=context.model,
device=device, device=device,
executor=executor, executor=executor,
model_fn=cfg.model_fn,
**strategy_kwargs, **strategy_kwargs,
) )
# Enable online rollout when the train_type is an ``online_*`` variant.
is_online = cfg.strategy.startswith("online_")
if is_online:
if not context.strategy.supports_online():
raise ValueError(
f"Strategy '{cfg.strategy}' does not support online rollout"
)
if cfg.reward_model_fn is None:
raise ValueError("reward_model_fn is required for online RL strategies")
tokenizer = AutoTokenizer.from_pretrained(self._param_path)
reward_model = cfg.reward_model_fn()
group_size = strategy_kwargs.get("group_size", 1)
rollout_batch_size = group_size * max(1, cfg.batch_per_device)
max_seq_len = getattr(context.model.config, "max_position_embeddings", None)
scheduler = InferenceScheduler(
model=context.model,
tokenizer=tokenizer,
max_batch_size=rollout_batch_size,
max_seq_len=max_seq_len,
max_prompt_len=max_seq_len or 4096,
)
generator = RolloutGenerator(
scheduler=scheduler,
tokenizer=tokenizer,
max_tokens=cfg.rollout_max_tokens,
group_size=group_size,
temperature=cfg.rollout_temperature,
top_k=cfg.rollout_top_k,
top_p=cfg.rollout_top_p,
)
runner = RolloutRunner(
generator=generator,
reward_model=reward_model,
rollout_interval=cfg.rollout_interval,
)
context.strategy.set_rollout_runner(runner)
return context return context
+21
View File
@@ -1,8 +1,14 @@
import logging import logging
from typing import List, Optional from typing import List, Optional
import torch.distributed as dist
from astrai.config import TrainConfig from astrai.config import TrainConfig
from astrai.parallel.setup import spawn_parallel_fn from astrai.parallel.setup import spawn_parallel_fn
from astrai.parallel.signal_handler import (
register_signal_handlers,
unregister_signal_handlers,
)
from astrai.trainer.train_callback import ( from astrai.trainer.train_callback import (
CallbackFactory, CallbackFactory,
TrainCallback, TrainCallback,
@@ -58,6 +64,7 @@ class Trainer:
.with_param_path(param_path, resume=resume) .with_param_path(param_path, resume=resume)
.build() .build()
) )
register_signal_handlers(context)
executor = context.executor executor = context.executor
self._call_callbacks("on_train_begin", context) self._call_callbacks("on_train_begin", context)
@@ -65,10 +72,14 @@ class Trainer:
context.model.train() context.model.train()
for epoch in range(context.epoch, context.config.n_epoch): for epoch in range(context.epoch, context.config.n_epoch):
if context.stop_requested:
break
context.epoch = epoch context.epoch = epoch
self._call_callbacks("on_epoch_begin", context) self._call_callbacks("on_epoch_begin", context)
for batch in context.dataloader: for batch in context.dataloader:
if context.stop_requested:
break
with executor.accumulate(context.model): with executor.accumulate(context.model):
self._call_callbacks("on_batch_begin", context) self._call_callbacks("on_batch_begin", context)
loss = context.strategy(batch) loss = context.strategy(batch)
@@ -83,6 +94,7 @@ class Trainer:
if executor.sync_gradients: if executor.sync_gradients:
self._call_callbacks("on_optimizer_step", context) self._call_callbacks("on_optimizer_step", context)
context.optimizer.step() context.optimizer.step()
context.strategy.on_optimizer_step()
context.optimizer.zero_grad() context.optimizer.zero_grad()
if context.scheduler: if context.scheduler:
@@ -90,12 +102,21 @@ class Trainer:
self._call_callbacks("on_epoch_end", context) self._call_callbacks("on_epoch_end", context)
if context.stop_requested:
logger.warning(
"Training interrupted by signal, saving emergency checkpoint..."
)
self._call_callbacks("on_error", context)
except Exception as e: except Exception as e:
logger.error("Training failed: %s", str(e), exc_info=True) logger.error("Training failed: %s", str(e), exc_info=True)
self._call_callbacks("on_error", context) self._call_callbacks("on_error", context)
raise raise
finally: finally:
self._call_callbacks("on_train_end", context) self._call_callbacks("on_train_end", context)
if executor.use_distributed and dist.is_initialized():
dist.barrier()
unregister_signal_handlers()
def train(self, param_path: Optional[str] = None, resume: bool = False): def train(self, param_path: Optional[str] = None, resume: bool = False):
cfg = self.train_config cfg = self.train_config
+2 -47
View File
@@ -1,51 +1,6 @@
#include "attn_decode_split_kv.cuh" #include "attn_dispatchers.cuh"
#include "attn_entry_utils.cuh" #include "attn_entry_utils.cuh"
#ifndef ASTRAI_NO_MMA
#include "attn_decode_split_kv_mma.cuh"
#endif
// Scalar fallback: one warp per query head, split-KV across grid.z.
static void launch_scalar_decode(AttentionParams<bf16>& p) {
int group_size = p.q_head / p.kv_head;
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
alloc_split_partials(p);
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
attn_decode_split_kv_kernel<<<dim3(p.batch * p.kv_head, 1, p.num_splits), dim3(32, group_size), smem>>>(p);
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
#ifndef ASTRAI_NO_MMA
// MMA head-packing requires G <= 16 (BR=16 rows). sm_80+ tensor-core
// + cp.async wins even at G=1 (decode is memory-bound, not compute-bound).
// STAGES=2 (double-buffer) for D<=128 (smem 16 KB); STAGES=1 for D=256
// (double-buffer would be 32 KB, near the 48 KB static cap — keep single
// to preserve occupancy).
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
static void launch_mma_decode(AttentionParams<bf16>& p) {
int tiles_total = (p.kv_len + BC - 1) / BC;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
alloc_split_partials(p);
attn_decode_split_kv_mma_kernel<HEAD_DIM, BC, STAGES><<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
#endif
template <int HEAD_DIM>
static void dispatch_decode(AttentionParams<bf16>& p) {
#ifndef ASTRAI_NO_MMA
int G = p.q_head / p.kv_head;
if (G >= 1 && G <= 16) {
launch_mma_decode<HEAD_DIM, 32>(p);
return;
}
#endif
launch_scalar_decode(p);
}
torch::Tensor attn_decode( torch::Tensor attn_decode(
torch::Tensor q, torch::Tensor q,
torch::Tensor k, torch::Tensor k,
@@ -60,11 +15,11 @@ torch::Tensor attn_decode(
TORCH_CHECK(p.q_len == 1, "Q seq_len must be 1"); TORCH_CHECK(p.q_len == 1, "Q seq_len must be 1");
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32"); TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
// O matches Q's original layout
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options()); auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
auto O_view = (layout == 1) ? O.transpose(1, 2) : O; auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
p.o = (bf16*)O_view.data_ptr(); 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);
return O; return O;
} }
+22 -25
View File
@@ -2,16 +2,10 @@
#include <cuda_bf16.h> #include <cuda_bf16.h>
#include <float.h> #include <float.h>
#include "attn_common.h" #include "attn_common.h"
#include "attn_warp_utils.cuh"
using bf16 = __nv_bfloat16;
constexpr int DC_CHUNK = 64; constexpr int DC_CHUNK = 64;
__device__ inline float warp_reduce_sum(float val) { template <int HEAD_DIM, bool IsCausal, bool HasMask>
for (int offset = 16; offset > 0; offset >>= 1)
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
return val;
}
__global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) { __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
int batch = blockIdx.x / p.kv_head; int batch = blockIdx.x / p.kv_head;
int kv_head = blockIdx.x % p.kv_head; int kv_head = blockIdx.x % p.kv_head;
@@ -48,7 +42,8 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
// Load K into shared memory (gather from strided global) // Load K into shared memory (gather from strided global)
int total = this_chunk * p.head_dim; int total = this_chunk * p.head_dim;
for (int i = threadIdx.y * 32 + lane; i < total; i += blockDim.x * blockDim.y) { for (int i = threadIdx.y * 32 + lane; i < total;
i += blockDim.x * blockDim.y) {
int s = i / p.head_dim; int s = i / p.head_dim;
int d_dim = i % p.head_dim; int d_dim = i % p.head_dim;
int kv_idx = chunk_start + s; int kv_idx = chunk_start + s;
@@ -60,24 +55,30 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
for (int s = 0; s < this_chunk; s++) { for (int s = 0; s < this_chunk; s++) {
float partial = 0.0f; float partial = 0.0f;
for (int i = 0; i < hd_per_thread; i++) for (int i = 0; i < hd_per_thread; i++)
partial += q_reg[i] * __bfloat162float(k_smem[s * p.head_dim + lane * hd_per_thread + i]); partial += q_reg[i] * __bfloat162float(
k_smem[s * p.head_dim + lane * hd_per_thread + i]);
partial = warp_reduce_sum(partial) * p.scale; partial = warp_reduce_sum(partial) * p.scale;
int kv_idx = chunk_start + s; int kv_idx = chunk_start + s;
if (p.use_mask && p.mask && !p.mask[mask_base + kv_idx]) if constexpr (HasMask) {
if (!p.mask[mask_base + kv_idx])
partial = -FLT_MAX; partial = -FLT_MAX;
if (p.causal_offset >= 0 && kv_idx > p.causal_offset) }
if constexpr (IsCausal) {
if (kv_idx > p.causal_offset)
partial = -FLT_MAX; partial = -FLT_MAX;
}
float new_m = fmaxf(m, partial); float new_m = fmaxf(m, partial);
float alpha = expf(m - new_m); float alpha = expf(m - new_m);
float beta = expf(partial - new_m); float beta = expf(partial - new_m);
d = d * alpha + beta; d = d * alpha + beta;
// V: stride-based read int v_off = kv_base + kv_idx * p.kv_stride_l
int v_off = kv_base + kv_idx * p.kv_stride_l + lane * hd_per_thread * p.kv_stride_d; + lane * hd_per_thread * p.kv_stride_d;
for (int i = 0; i < hd_per_thread; i++) for (int i = 0; i < hd_per_thread; i++)
acc_reg[i] = acc_reg[i] * alpha + __bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta; acc_reg[i] = fmaf(acc_reg[i], alpha,
__bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta);
m = new_m; m = new_m;
} }
__syncthreads(); __syncthreads();
@@ -85,7 +86,7 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
// ---- write UN-normalised partials for this split ---- // ---- write UN-normalised partials for this split ----
size_t bh = (size_t)batch * p.q_head + q_head; size_t bh = (size_t)batch * p.q_head + q_head;
size_t slot = bh * p.num_splits + split; size_t slot = bh * MAX_SPLITS + split;
int d0 = lane * hd_per_thread; int d0 = lane * hd_per_thread;
for (int i = 0; i < hd_per_thread; i++) { for (int i = 0; i < hd_per_thread; i++) {
int dd = d0 + i; int dd = d0 + i;
@@ -97,9 +98,6 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
} }
} }
// Reduce split-K partials into the final bf16 output. One block per (batch,
// q_head); each thread folds across all splits with a single-pass
// online-rescale reduction (expf + FMA counts halved vs 3-pass original).
__global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) { __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
int bh = blockIdx.x; int bh = blockIdx.x;
int d = threadIdx.x; int d = threadIdx.x;
@@ -108,7 +106,7 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
int batch = bh / p.q_head; int batch = bh / p.q_head;
int q_head = bh % p.q_head; int q_head = bh % p.q_head;
size_t split_base = (size_t)bh * p.num_splits; size_t split_base = (size_t)bh * MAX_SPLITS;
const float* mlp = p.ml_part + split_base * 2; const float* mlp = p.ml_part + split_base * 2;
const float* op = p.o_part + split_base * p.head_dim; const float* op = p.o_part + split_base * p.head_dim;
@@ -118,15 +116,14 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
if (mi <= -FLT_MAX) continue; if (mi <= -FLT_MAX) continue;
float li = mlp[s * 2 + 1]; float li = mlp[s * 2 + 1];
float nm = fmaxf(m, mi); float nm = fmaxf(m, mi);
float corr = __expf(m - nm); float corr = expf(m - nm);
float e = __expf(mi - nm); float e = expf(mi - nm);
acc = acc * corr + op[s * p.head_dim + d] * e; acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
l = l * corr + li * e; l = fmaf(l, corr, li * e);
m = nm; m = nm;
} }
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f; float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
// Stride-based output write (q_len=1 for decode, so stride_l not needed)
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h + d * p.q_stride_d; int o_off = batch * p.q_stride_b + q_head * p.q_stride_h + d * p.q_stride_d;
p.o[o_off] = __float2bfloat16(acc * inv); p.o[o_off] = __float2bfloat16(acc * inv);
} }
+55 -70
View File
@@ -3,85 +3,72 @@
#include <cuda_bf16.h> #include <cuda_bf16.h>
#include "attn_common.h" #include "attn_common.h"
#include "attn_mma_utils.cuh" #include "attn_mma_utils.cuh"
#include "attn_warp_utils.cuh"
using bf16 = __nv_bfloat16;
// Split-K (FlashDecoding) tensor-core decode via GQA head-packing. // Split-K (FlashDecoding) tensor-core decode via GQA head-packing.
// Decode has q_len == 1, so we pack G = q_head/kv_head query heads into the
// M=16 rows of mma.sync.m16n8k16, turning G independent GEMVs into a single
// GEMM that reuses each loaded K/V tile across all G heads.
// //
// Decode has q_len == 1, so S = q @ K^T is a GEMV per head — no tensor-core // IsCausal and HasMask are compile-time bools — no runtime branch in the
// work on its own. But GQA gives us G = q_head / kv_head query heads that all // inner compute loop.
// share one kv_head. We pack those G heads into the M=16 rows of //
// mma.sync.m16n8k16, turning G independent GEMVs into a single GEMM that // Traits = KernelTraits<HEAD_DIM, BC=32, WARPS=1, STAGES=<2 or 1>>.
// reuses each loaded K/V tile across all G heads (K/V load is the decode template <typename Traits, bool IsCausal, bool HasMask>
// bottleneck, so the reuse is the win, not the flops). The KV sequence is
// partitioned across gridDim.z blocks so that a decode with only
// batch*kv_head independent tasks can fill all SMs. Each (batch, kv_head,
// split) block computes an UN-normalised partial (Oacc, m, l) over its KV
// slice; the combine kernel below reduces across splits. Fixes the "grid too
// small" bottleneck (0.04 waves/SM → many blocks) for long-context,
// small-batch decode.
template <int HEAD_DIM, int BC, int STAGES = 2>
__global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) { __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
constexpr int KD = HEAD_DIM / 16;
constexpr int NC8 = BC / 8;
constexpr int KT2 = BC / 16;
constexpr int DN8 = HEAD_DIM / 8;
constexpr int LD = HEAD_DIM;
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
constexpr int VEC = 8;
constexpr int TOTAL = BC * HEAD_DIM;
const int lane = threadIdx.x; const int lane = threadIdx.x;
const int gid = lane >> 2; const int gid = lane >> 2;
const int tid4 = lane & 3; const int tid4 = lane & 3;
const int kv_head = blockIdx.x; const int pass = blockIdx.x / p.kv_head;
const int kv_head = blockIdx.x % p.kv_head;
const int batch = blockIdx.y; const int batch = blockIdx.y;
const int split = blockIdx.z; const int split = blockIdx.z;
const int G = p.q_head / p.kv_head;
const int q_head0 = kv_head * G;
// Double-buffered shared memory for K/V (no sQ needed — Q goes direct constexpr int MAX_G = 16;
// from global to registers). const int G_total = p.q_head / p.kv_head;
__shared__ __align__(16) bf16 sK[STAGES * BC * LD]; const int g_begin = pass * MAX_G;
__shared__ __align__(16) bf16 sV[STAGES * BC * LD]; const int G = min(MAX_G, G_total - g_begin);
const int q_head0 = kv_head * G_total + g_begin;
// ---- Load Q directly from global into mma A-operand registers ---- // Double-buffered shared memory for K/V (no sQ needed)
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
// Load Q directly from global into mma A-operand registers.
// stride_row = p.q_stride_h for decode (q_len=1).
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h; const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
const int qra = gid; const int qra = gid;
const int qrb = gid + 8; const int qrb = gid + 8;
const bool va = qra < G, vb = qrb < G; const bool va = qra < G, vb = qrb < G;
unsigned Qa[KD][4]; unsigned Qa[Traits::KD][4];
load_q_mma_frags<KD>(p.q + q_base, p.q_stride_h, p.q_stride_d, load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
qra, qrb, va, vb, tid4, Qa); qra, qrb, va, vb, tid4, Qa);
float Oacc[DN8][4]; float Oacc[Traits::DN8][4];
#pragma unroll #pragma unroll
for (int j = 0; j < DN8; j++) for (int j = 0; j < Traits::DN8; j++)
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f; Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f; float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
// KV: stride-based base — [batch, kv_head, kv_len, head_dim]
const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h; const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
const int tiles_total = (p.kv_len + BC - 1) / BC; const int tiles_total = (p.kv_len + Traits::BC - 1) / Traits::BC;
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits; const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
const int ti_begin = split * tiles_per_split; const int ti_begin = split * tiles_per_split;
const int ti_end = min(tiles_total, ti_begin + tiles_per_split); const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
const int has_mask = p.use_mask && p.mask;
// ---- Load tile lambda: predicated cp.async, unified full/partial ---- // ---- Load tile lambda: predicated cp.async ----
auto load_tile = [&](int ti, int buf) { auto load_tile = [&](int ti, int buf) {
int kv0 = ti * BC; int kv0 = ti * Traits::BC;
bf16* dK = sK + buf * BC * LD; bf16* dK = sK + buf * Traits::BC * Traits::LD;
bf16* dV = sV + buf * BC * LD; bf16* dV = sV + buf * Traits::BC * Traits::LD;
#pragma unroll #pragma unroll
for (int i = lane * VEC; i < TOTAL; i += 32 * VEC) { for (int i = lane * Traits::VEC; i < Traits::TOTAL;
int r = i / HEAD_DIM, d = i % HEAD_DIM; i += Traits::NUM_THREADS * Traits::VEC) {
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
int kc = kv0 + r; int kc = kv0 + r;
bool valid = kc < p.kv_len; bool valid = kc < p.kv_len;
int off = r * LD + swiz_col(d, r, SWIZ_MASK); int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
// KV stride-based: contiguous within head_dim (stride_d == 1 typically)
int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d; int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d;
cp_async_16_pred(&dK[off], &p.k[g_off], valid); cp_async_16_pred(&dK[off], &p.k[g_off], valid);
cp_async_16_pred(&dV[off], &p.v[g_off], valid); cp_async_16_pred(&dV[off], &p.v[g_off], valid);
@@ -89,50 +76,48 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
cp_async_commit(); cp_async_commit();
}; };
// ---- Prologue: issue first tile load ---- constexpr int BUF_MASK = (Traits::STAGES > 1) ? (Traits::STAGES - 1) : 0;
// Prologue
if (ti_begin < ti_end) { if (ti_begin < ti_end) {
load_tile(ti_begin, 0); load_tile(ti_begin, 0);
} }
for (int ti = ti_begin; ti < ti_end; ti++) { for (int ti = ti_begin; ti < ti_end; ti++) {
constexpr int BUF_MASK = (STAGES > 1) ? (STAGES - 1) : 0;
int buf = (ti - ti_begin) & BUF_MASK; int buf = (ti - ti_begin) & BUF_MASK;
// Wait for current tile, then issue next tile's prefetch (overlaps
// with this tile's compute). Single syncwarp covers both hazards.
// When STAGES==1, no prefetch — load happens at end of prior iter.
cp_async_wait_group<0>(); cp_async_wait_group<0>();
__syncwarp(); __syncwarp();
if constexpr (STAGES > 1) { if constexpr (Traits::STAGES > 1) {
if (ti + 1 < ti_end) if (ti + 1 < ti_end)
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK); load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
} }
const bf16* bK = sK + buf * BC * LD; const bf16* bK = sK + buf * Traits::BC * Traits::LD;
const bf16* bV = sV + buf * BC * LD; const bf16* bV = sV + buf * Traits::BC * Traits::LD;
int kv0 = ti * BC; int kv0 = ti * Traits::BC;
float Sacc[NC8][4]; float Sacc[Traits::NC8][4];
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc); mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
#pragma unroll #pragma unroll
for (int n8 = 0; n8 < NC8; n8++) for (int n8 = 0; n8 < Traits::NC8; n8++)
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale, Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale; Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
// Decode: q_len=1, so qrow0=qrow1=0, mask_q_stride irrelevant // Decode: q_len=1, so qrow0=qrow1=0
int maxc = (p.causal_offset >= 0) ? min(p.kv_len, p.causal_offset + 1) : p.kv_len; int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
mma_softmax_tile<NC8, DN8>(kv0, maxc, maxc, mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
0, 0, 0, 0,
p.mask_b_stride, 0, p.mask_b_stride, 0,
batch, batch,
p.mask, has_mask, p.mask,
Sacc, Oacc, m0, m1, l0, l1, lane); Sacc, Oacc, m0, m1, l0, l1, lane);
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc); mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
__syncwarp(); __syncwarp();
if constexpr (STAGES == 1) { if constexpr (Traits::STAGES == 1) {
if (ti + 1 < ti_end) if (ti + 1 < ti_end)
load_tile(ti + 1, 0); load_tile(ti + 1, 0);
} }
@@ -141,21 +126,21 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
// ---- write UN-normalised partials for this split ---- // ---- write UN-normalised partials for this split ----
auto split_slot = [&](int h) -> size_t { auto split_slot = [&](int h) -> size_t {
size_t bh = (size_t)batch * p.q_head + h; size_t bh = (size_t)batch * p.q_head + h;
return bh * p.num_splits + split; return bh * MAX_SPLITS + split;
}; };
#pragma unroll #pragma unroll
for (int dn8 = 0; dn8 < DN8; dn8++) { for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
int d = dn8 * 8 + 2 * tid4; int d = dn8 * 8 + 2 * tid4;
int r0 = gid, r1 = gid + 8; int r0 = gid, r1 = gid + 8;
if (r0 < G) { if (r0 < G) {
int h = q_head0 + r0; int h = q_head0 + r0;
float* op = p.o_part + split_slot(h) * HEAD_DIM; float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
op[d] = Oacc[dn8][0]; op[d] = Oacc[dn8][0];
op[d + 1] = Oacc[dn8][1]; op[d + 1] = Oacc[dn8][1];
} }
if (r1 < G) { if (r1 < G) {
int h = q_head0 + r1; int h = q_head0 + r1;
float* op = p.o_part + split_slot(h) * HEAD_DIM; float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
op[d] = Oacc[dn8][2]; op[d] = Oacc[dn8][2];
op[d + 1] = Oacc[dn8][3]; op[d + 1] = Oacc[dn8][3];
} }
+195
View File
@@ -0,0 +1,195 @@
#pragma once
// Shared attention dispatchers — used by both production .cu and test .cu.
// No torch dependency; pure CUDA.
#include <cuda_runtime.h>
#include <algorithm>
#include "attn_warp_utils.cuh"
#include "attn_prefill_split_q.cuh"
#include "attn_decode_split_kv.cuh"
#include "attn_paged_decode_split_kv.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"
#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)));
}
// ======================================================================
// Prefill
// ======================================================================
#ifndef ASTRAI_NO_MMA
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static inline void launch_prefill_mma(AttentionParams<bf16>& p) {
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);
}
#endif
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static inline void launch_prefill_scalar(AttentionParams<bf16>& p) {
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);
}
template <int HEAD_DIM>
static inline void dispatch_prefill(AttentionParams<bf16>& p) {
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);
}
#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);
}
#endif
}
// ======================================================================
// Decode
// ======================================================================
#ifndef ASTRAI_NO_MMA
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static inline void launch_decode_mma(AttentionParams<bf16>& p, int group_size) {
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>;
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask><<<grid, 32>>>(p);
}
#endif
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static inline void launch_decode_scalar(AttentionParams<bf16>& p, int group_size) {
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);
}
template <int HEAD_DIM>
static inline void dispatch_decode(AttentionParams<bf16>& p) {
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);
}
#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);
}
#endif
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
// ======================================================================
// Paged Decode
// ======================================================================
#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) {
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);
}
}
#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;
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
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);
}
template <int HEAD_DIM>
static inline void dispatch_paged_decode(PagedAttentionParams<bf16>& p) {
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);
}
#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);
}
#endif
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
+3 -9
View File
@@ -2,16 +2,10 @@
#include <torch/extension.h> #include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h> #include <c10/cuda/CUDAGuard.h>
#include "attn_common.h" #include "attn_common.h"
#include "attn_warp_utils.cuh"
using bf16 = __nv_bfloat16; using bf16 = __nv_bfloat16;
inline int compute_num_splits(int base_blocks, int tiles_total) {
int sm_count = 0;
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
return std::max(1, std::min(n, std::min(tiles_total, 32)));
}
// Dispatch head_dim: shared macro — avoids C++20 lambda template syntax. // Dispatch head_dim: shared macro — avoids C++20 lambda template syntax.
// Usage: DISPATCH_HEAD_DIM(hd, fn, arg) // Usage: DISPATCH_HEAD_DIM(hd, fn, arg)
// Expands to: fn<32>(arg); fn<64>(arg); etc. // Expands to: fn<32>(arg); fn<64>(arg); etc.
@@ -29,8 +23,8 @@ inline int compute_num_splits(int base_blocks, int tiles_total) {
template<typename P> template<typename P>
inline void alloc_split_partials(P& p) { inline void alloc_split_partials(P& p) {
auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA); auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
auto o_part = torch::empty({p.batch, p.q_head, p.num_splits, p.head_dim}, fopt); auto o_part = torch::empty({p.batch, p.q_head, MAX_SPLITS, p.head_dim}, fopt);
auto ml_part = torch::empty({p.batch, p.q_head, p.num_splits, 2}, fopt); auto ml_part = torch::empty({p.batch, p.q_head, MAX_SPLITS, 2}, fopt);
p.o_part = (float*)o_part.data_ptr(); p.o_part = (float*)o_part.data_ptr();
p.ml_part = (float*)ml_part.data_ptr(); p.ml_part = (float*)ml_part.data_ptr();
} }
+75 -77
View File
@@ -3,10 +3,41 @@
#include <cuda_fp16.h> #include <cuda_fp16.h>
#include <cuda_runtime.h> #include <cuda_runtime.h>
// Shared MMA utilities for tensor-core GQA kernels. // ============================================================================
// mma.sync.m16n8k16 PTX wrappers, ldmatrix helpers, and bf16 packing. // KernelTraits — FlashAttention-v2 style compile-time configuration bundle.
//
// Bundles all dimension-dependent constants so device functions only need a
// single Traits template parameter rather than scattered <KD, NC8, KT2, ...>.
// ============================================================================
template <int HEAD_DIM_, int BC_, int WARPS_, int STAGES_>
struct KernelTraits {
static constexpr int HEAD_DIM = HEAD_DIM_;
static constexpr int BC = BC_; // K/V tile size along seq dim
static constexpr int WARPS = WARPS_; // warps per block
static constexpr int STAGES = STAGES_; // double-buffer stages (1 or 2)
static constexpr int BR = 16; // Q rows per warp (mma M=16)
// Derived: mma.sync.m16n8k16 tile counts
static constexpr int KD = HEAD_DIM / 16; // Q/K k-slides
static constexpr int NC8 = BC / 8; // S n-tiles (N=8)
static constexpr int KT2 = BC / 16; // P k-tiles (K=16)
static constexpr int DN8 = HEAD_DIM / 8; // O n-tiles (N=8)
static constexpr int LD = HEAD_DIM; // smem leading dim
// XOR swizzle chunk bits for ldmatrix bank-conflict avoidance.
// mask = log2(LD/8) bits, clamped to stay within LD.
static constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
static constexpr int NUM_THREADS = WARPS * 32;
static constexpr int VEC = 8; // bf16 per cp.async unit (16 bytes)
static constexpr int TOTAL = BC * HEAD_DIM; // total elements per tile
};
// ---- PTX wrappers ----
using bf16 = __nv_bfloat16;
// mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32
__device__ __forceinline__ void mma16816(float* d, const unsigned* a, __device__ __forceinline__ void mma16816(float* d, const unsigned* a,
const unsigned* b, const float* c) { const unsigned* b, const float* c) {
asm volatile( asm volatile(
@@ -37,9 +68,7 @@ __device__ __forceinline__ unsigned pkb(bf16 a, bf16 b) {
} }
// ldmatrix: cooperatively load mma fragments from smem (one instruction per // ldmatrix: cooperatively load mma fragments from smem (one instruction per
// 16x16 / 16x8 tile) with the exact register layout mma expects — replaces the // 16x16 / 16x8 tile) with the exact register layout mma expects.
// scalar per-thread fragment packing, cutting shared-load instructions and bank
// conflicts. Each lane supplies the shared address of one 8-wide row.
__device__ __forceinline__ void ldmatrix_x4(unsigned* r, const bf16* p) { __device__ __forceinline__ void ldmatrix_x4(unsigned* r, const bf16* p) {
unsigned a = __cvta_generic_to_shared(p); unsigned a = __cvta_generic_to_shared(p);
asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];" asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];"
@@ -60,29 +89,19 @@ __device__ __forceinline__ void ldmatrix_x2_trans(unsigned* r, const bf16* p) {
} }
// XOR swizzle for shared-memory column at 8-bf16 chunk granularity. // XOR swizzle for shared-memory column at 8-bf16 chunk granularity.
// Eliminates ldmatrix bank conflicts without LD padding: consecutive rows
// land in distinct bank groups. swiz_col(d, r, mask) = ((d>>3)^(r&mask))<<3 | (d&7).
// mask must cover log2(HEAD_DIM/8) chunk bits but stay within LD: use 7 for
// HEAD_DIM>=64 (8+ chunks), 3 for HEAD_DIM=32 (4 chunks). Default 7 keeps
// existing HEAD_DIM>=64 call sites working unchanged.
__device__ __forceinline__ int swiz_col(int d, int r, int mask = 7) { __device__ __forceinline__ int swiz_col(int d, int r, int mask = 7) {
return ((d >> 3) ^ (r & mask)) << 3 | (d & 7); return ((d >> 3) ^ (r & mask)) << 3 | (d & 7);
} }
// cp.async: copy 16 bytes (8 bf16) from global to shared memory directly, // cp.async: copy 16 bytes (8 bf16) from global to shared memory directly.
// bypassing registers. Eliminates shared-store bank conflicts and cuts
// load-loop instruction count in half (1 cp.async vs 1 LDG + 1 STS).
// Requires sm_80+.
__device__ __forceinline__ void cp_async_16(bf16* smem_ptr, const void* gmem_ptr) { __device__ __forceinline__ void cp_async_16(bf16* smem_ptr, const void* gmem_ptr) {
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr); unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
asm volatile("cp.async.ca.shared.global [%0], [%1], 16;" asm volatile("cp.async.ca.shared.global [%0], [%1], 16;"
:: "r"(smem_addr), "l"(gmem_ptr)); :: "r"(smem_addr), "l"(gmem_ptr));
} }
// Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill the // Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill.
// destination (src-size operand = 0 → no bytes read from src, so an // src_size=0 → no bytes read from src, so out-of-bounds src address is safe.
// out-of-bounds src address is never dereferenced). Lets full and partial
// tiles share one uniform async load path — no scalar fallback branch.
__device__ __forceinline__ void cp_async_16_pred(bf16* smem_ptr, __device__ __forceinline__ void cp_async_16_pred(bf16* smem_ptr,
const void* gmem_ptr, const void* gmem_ptr,
bool pred) { bool pred) {
@@ -100,9 +119,6 @@ __device__ __forceinline__ void cp_async_wait_all() {
asm volatile("cp.async.wait_all;"); asm volatile("cp.async.wait_all;");
} }
// Wait until at most N commit groups are still in flight. Used for
// double-buffered pipelining: wait_group<1> lets the next tile's cp.async
// continue while ensuring the current tile's data is ready.
template <int N> template <int N>
__device__ __forceinline__ void cp_async_wait_group() { __device__ __forceinline__ void cp_async_wait_group() {
asm volatile("cp.async.wait_group %0;" :: "n"(N)); asm volatile("cp.async.wait_group %0;" :: "n"(N));
@@ -139,78 +155,65 @@ __device__ inline void load_q_mma_frags(
} }
// --------------------------------------------------------------------------- // ---------------------------------------------------------------------------
// Shared MMA compute functions — used by both decode and prefill MMA kernels.
// Extracted because S=Q@K^T, online softmax, and P@V are structurally identical
// between the two kernels; only the per-row causal/mask bounds differ.
// ---------------------------------------------------------------------------
// S = Q @ K^T (Qa pre-loaded by the caller; scale applied post-mma in the // S = Q @ K^T (Qa pre-loaded by the caller; scale applied post-mma in the
// caller to avoid bf16 precision loss). // caller to avoid bf16 precision loss).
// LD and SWIZ_MASK are constexpr in the calling kernel — passing them as // Traits provides KD, NC8, LD, and SWIZ_MASK.
// runtime ints lets the compiler fold them while keeping the signature clean. // ---------------------------------------------------------------------------
template <int KD, int NC8> template <typename Traits>
__device__ inline void mma_compute_scores( __device__ inline void mma_compute_scores(
const unsigned Qa[KD][4], const unsigned Qa[Traits::KD][4],
const bf16* __restrict__ sK, const bf16* __restrict__ sK,
int LD,
int SWIZ_MASK,
int lane, int lane,
float Sacc[NC8][4]) float Sacc[Traits::NC8][4])
{ {
#pragma unroll #pragma unroll
for (int n8 = 0; n8 < NC8; n8++) { for (int n8 = 0; n8 < Traits::NC8; n8++) {
Sacc[n8][0] = Sacc[n8][1] = Sacc[n8][2] = Sacc[n8][3] = 0.0f; Sacc[n8][0] = Sacc[n8][1] = Sacc[n8][2] = Sacc[n8][3] = 0.0f;
int krow_l = n8 * 8 + (lane & 7); int krow_l = n8 * 8 + (lane & 7);
int kcol_h = (lane & 8) ? 8 : 0; int kcol_h = (lane & 8) ? 8 : 0;
#pragma unroll #pragma unroll
for (int kt = 0; kt < KD; kt++) { for (int kt = 0; kt < Traits::KD; kt++) {
unsigned b[2]; unsigned b[2];
ldmatrix_x2(b, &sK[krow_l * LD + swiz_col(kt * 16 + kcol_h, krow_l, SWIZ_MASK)]); ldmatrix_x2(b, &sK[krow_l * Traits::LD
+ swiz_col(kt * 16 + kcol_h, krow_l, Traits::SWIZ_MASK)]);
mma16816(Sacc[n8], Qa[kt], b, Sacc[n8]); mma16816(Sacc[n8], Qa[kt], b, Sacc[n8]);
} }
} }
} }
// ---------------------------------------------------------------------------
// Online softmax + Oacc rescale for one K/V tile. // Online softmax + Oacc rescale for one K/V tile.
// maxc0/maxc1: per-row KV column bounds (prefill: per-query-row causal limits; //
// decode: same value for both rows since q_len==1). // HasMask is a compile-time template bool: when false, the mask branch is
// qrow0/qrow1: query row indices (for 3D mask indexing; decode passes 0). // entirely dead-code-eliminated from the inner unrolled loop.
// mask_b_stride/mask_q_stride: mask layout (2D: mask_q_stride=0; 3D: =kv_len). // ---------------------------------------------------------------------------
// Reads Sacc (Q@K^T scores), applies causal/mask, computes P = exp(S - nm), template <typename Traits, bool HasMask>
// rescales Oacc by exp(m_old - nm), and updates m/l — all in place.
template <int NC8, int DN8>
__device__ inline void mma_softmax_tile( __device__ inline void mma_softmax_tile(
int kv0, int kv0,
int maxc0, int maxc0, int maxc1,
int maxc1, int qrow0, int qrow1,
int qrow0, int mask_b_stride, int mask_q_stride,
int qrow1,
int mask_b_stride,
int mask_q_stride,
int mask_batch, int mask_batch,
const bool* __restrict__ mask, const bool* __restrict__ mask,
bool has_mask, float Sacc[Traits::NC8][4],
float Sacc[NC8][4], float Oacc[Traits::DN8][4],
float Oacc[DN8][4],
float& m0, float& m1, float& m0, float& m1,
float& l0, float& l1, float& l0, float& l1,
int lane) int lane)
{ {
int tid4 = lane & 3; int tid4 = lane & 3;
// Mask out-of-bounds / masked columns: set -FLT_MAX so expf → 0 downstream
// without per-element sentinel checks. Compute tile-local row maxima.
float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX; float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX;
int mask_base0 = mask_batch * mask_b_stride + qrow0 * mask_q_stride; int mask_base0 = mask_batch * mask_b_stride + qrow0 * mask_q_stride;
int mask_base1 = mask_batch * mask_b_stride + qrow1 * mask_q_stride; int mask_base1 = mask_batch * mask_b_stride + qrow1 * mask_q_stride;
#pragma unroll #pragma unroll
for (int n8 = 0; n8 < NC8; n8++) { for (int n8 = 0; n8 < Traits::NC8; n8++) {
int cc = kv0 + n8 * 8 + 2 * tid4; int cc = kv0 + n8 * 8 + 2 * tid4;
int c1 = cc + 1; int c1 = cc + 1;
bool b0 = (cc >= maxc0) || (has_mask && !mask[mask_base0 + cc]); bool b0 = (cc >= maxc0) || (HasMask && !mask[mask_base0 + cc]);
bool b1 = (c1 >= maxc0) || (has_mask && !mask[mask_base0 + c1]); bool b1 = (c1 >= maxc0) || (HasMask && !mask[mask_base0 + c1]);
bool b2 = (cc >= maxc1) || (has_mask && !mask[mask_base1 + cc]); bool b2 = (cc >= maxc1) || (HasMask && !mask[mask_base1 + cc]);
bool b3 = (c1 >= maxc1) || (has_mask && !mask[mask_base1 + c1]); bool b3 = (c1 >= maxc1) || (HasMask && !mask[mask_base1 + c1]);
float s0 = b0 ? -FLT_MAX : Sacc[n8][0]; float s0 = b0 ? -FLT_MAX : Sacc[n8][0];
float s1 = b1 ? -FLT_MAX : Sacc[n8][1]; float s1 = b1 ? -FLT_MAX : Sacc[n8][1];
float s2 = b2 ? -FLT_MAX : Sacc[n8][2]; float s2 = b2 ? -FLT_MAX : Sacc[n8][2];
@@ -220,29 +223,20 @@ __device__ inline void mma_softmax_tile(
rmax0 = fmaxf(rmax0, fmaxf(s0, s1)); rmax0 = fmaxf(rmax0, fmaxf(s0, s1));
rmax1 = fmaxf(rmax1, fmaxf(s2, s3)); rmax1 = fmaxf(rmax1, fmaxf(s2, s3));
} }
// Warp-reduce row maxima across the 4-lane thread group (xor 1, xor 2).
rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 1)); rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 1));
rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 2)); rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 2));
rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 1)); rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 1));
rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 2)); rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 2));
// nm = max(running max m, tile-local max rmax) — updated running maximum.
float nm0 = fmaxf(m0, rmax0), nm1 = fmaxf(m1, rmax1); float nm0 = fmaxf(m0, rmax0), nm1 = fmaxf(m1, rmax1);
// corr rescales Oacc and l by exp(m_old - nm). When all-masked (m == nm ==
// -FLT_MAX), exp(0) = 1 — correct, no guard needed.
float corr0 = __expf(m0 - nm0); float corr0 = __expf(m0 - nm0);
float corr1 = __expf(m1 - nm1); float corr1 = __expf(m1 - nm1);
// pn guards only the all-masked-row edge: if nm == -FLT_MAX, exp(S - nm)
// gives 1 not 0 for masked entries. Two scalar masks replace 4*NC8
// per-element comparisons.
float pn0 = (nm0 == -FLT_MAX) ? 0.0f : 1.0f; float pn0 = (nm0 == -FLT_MAX) ? 0.0f : 1.0f;
float pn1 = (nm1 == -FLT_MAX) ? 0.0f : 1.0f; float pn1 = (nm1 == -FLT_MAX) ? 0.0f : 1.0f;
// P = exp(S - nm) for each element. Masked entries (Sacc = -FLT_MAX) give
// exp(-inf) ≈ 0 naturally; pn zero-fills the all-masked-row edge.
float rsum0 = 0.0f, rsum1 = 0.0f; float rsum0 = 0.0f, rsum1 = 0.0f;
#pragma unroll #pragma unroll
for (int n8 = 0; n8 < NC8; n8++) { for (int n8 = 0; n8 < Traits::NC8; n8++) {
float p0 = pn0 * __expf(Sacc[n8][0] - nm0); float p0 = pn0 * __expf(Sacc[n8][0] - nm0);
float p1 = pn0 * __expf(Sacc[n8][1] - nm0); float p1 = pn0 * __expf(Sacc[n8][1] - nm0);
float p2 = pn1 * __expf(Sacc[n8][2] - nm1); float p2 = pn1 * __expf(Sacc[n8][2] - nm1);
@@ -261,22 +255,25 @@ __device__ inline void mma_softmax_tile(
m0 = nm0; m1 = nm1; m0 = nm0; m1 = nm1;
#pragma unroll #pragma unroll
for (int j = 0; j < DN8; j++) { for (int j = 0; j < Traits::DN8; j++) {
Oacc[j][0] *= corr0; Oacc[j][1] *= corr0; Oacc[j][0] *= corr0; Oacc[j][1] *= corr0;
Oacc[j][2] *= corr1; Oacc[j][3] *= corr1; Oacc[j][2] *= corr1; Oacc[j][3] *= corr1;
} }
} }
// ---------------------------------------------------------------------------
// O += P @ V (Sacc must contain P = attention weights after softmax). // O += P @ V (Sacc must contain P = attention weights after softmax).
template <int DN8, int KT2> // Traits provides DN8, KT2, LD, and SWIZ_MASK.
// ---------------------------------------------------------------------------
template <typename Traits>
__device__ inline void mma_pv_accumulate( __device__ inline void mma_pv_accumulate(
float Sacc[][4], float Sacc[][4],
const bf16* __restrict__ sV, const bf16* __restrict__ sV,
int LD, int SWIZ_MASK, int lane, int lane,
float Oacc[DN8][4]) float Oacc[Traits::DN8][4])
{ {
#pragma unroll #pragma unroll
for (int kt2 = 0; kt2 < KT2; kt2++) { for (int kt2 = 0; kt2 < Traits::KT2; kt2++) {
unsigned Pa[4]; unsigned Pa[4];
Pa[0] = pk2(Sacc[kt2 * 2][0], Sacc[kt2 * 2][1]); Pa[0] = pk2(Sacc[kt2 * 2][0], Sacc[kt2 * 2][1]);
Pa[1] = pk2(Sacc[kt2 * 2][2], Sacc[kt2 * 2][3]); Pa[1] = pk2(Sacc[kt2 * 2][2], Sacc[kt2 * 2][3]);
@@ -284,9 +281,10 @@ __device__ inline void mma_pv_accumulate(
Pa[3] = pk2(Sacc[kt2 * 2 + 1][2], Sacc[kt2 * 2 + 1][3]); Pa[3] = pk2(Sacc[kt2 * 2 + 1][2], Sacc[kt2 * 2 + 1][3]);
int vrow_l = kt2 * 16 + (lane & 15); int vrow_l = kt2 * 16 + (lane & 15);
#pragma unroll #pragma unroll
for (int dn8 = 0; dn8 < DN8; dn8++) { for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
unsigned b[2]; unsigned b[2];
ldmatrix_x2_trans(b, &sV[vrow_l * LD + swiz_col(dn8 * 8, vrow_l, SWIZ_MASK)]); ldmatrix_x2_trans(b, &sV[vrow_l * Traits::LD
+ swiz_col(dn8 * 8, vrow_l, Traits::SWIZ_MASK)]);
mma16816(Oacc[dn8], Pa, b, Oacc[dn8]); mma16816(Oacc[dn8], Pa, b, Oacc[dn8]);
} }
} }
+2 -42
View File
@@ -1,47 +1,6 @@
#include "attn_paged_decode_split_kv.cuh" #include "attn_dispatchers.cuh"
#ifndef ASTRAI_NO_MMA
#include "attn_paged_decode_split_kv_mma.cuh"
#endif
#include "attn_entry_utils.cuh" #include "attn_entry_utils.cuh"
static void launch_paged_scalar_decode(PagedAttentionParams<bf16>& p) {
int group_size = p.q_head / p.kv_head;
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
alloc_split_partials(p);
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
dim3 grid = dim3(p.batch * p.kv_head, 1, p.num_splits);
dim3 block = dim3(32, group_size);
paged_attn_decode_split_kv_kernel<<<grid, block, smem>>>(p);
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
#ifndef ASTRAI_NO_MMA
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
static void launch_paged_mma_decode(PagedAttentionParams<bf16>& p) {
int tiles_total = (p.kv_len + BC - 1) / BC;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
alloc_split_partials(p);
paged_attn_decode_split_kv_mma_kernel<HEAD_DIM, BC, STAGES><<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
#endif
template <int HEAD_DIM>
static void dispatch_paged_decode(PagedAttentionParams<bf16>& p) {
#ifndef ASTRAI_NO_MMA
int G = p.q_head / p.kv_head;
if (G >= 1 && G <= 16 && p.page_size >= 32) {
launch_paged_mma_decode<HEAD_DIM, 32>(p);
return;
}
#endif
launch_paged_scalar_decode(p);
}
torch::Tensor attn_paged_decode( torch::Tensor attn_paged_decode(
torch::Tensor q, torch::Tensor q,
torch::Tensor page_table, torch::Tensor page_table,
@@ -62,6 +21,7 @@ torch::Tensor attn_paged_decode(
auto O_view = (layout == 1) ? O.transpose(1, 2) : O; auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
p.o = (bf16*)O_view.data_ptr(); p.o = (bf16*)O_view.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);
return O; return O;
} }
+22 -23
View File
@@ -2,17 +2,10 @@
#include <cuda_bf16.h> #include <cuda_bf16.h>
#include <float.h> #include <float.h>
#include "attn_common.h" #include "attn_common.h"
#include "attn_warp_utils.cuh"
using bf16 = __nv_bfloat16;
constexpr int PDC_CHUNK = 64; constexpr int PDC_CHUNK = 64;
__device__ inline float paged_warp_reduce_sum(float val) { template <int HEAD_DIM, bool IsCausal, bool HasMask>
for (int offset = 16; offset > 0; offset >>= 1)
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
return val;
}
// Split-KV scalar decode: one warp per query head, grid.z partitions KV.
__global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p) { __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p) {
int batch = blockIdx.x / p.kv_head; int batch = blockIdx.x / p.kv_head;
int kv_head = blockIdx.x % p.kv_head; int kv_head = blockIdx.x % p.kv_head;
@@ -22,7 +15,6 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
int lane = threadIdx.x; int lane = threadIdx.x;
int hd_per_thread = p.head_dim / 32; int hd_per_thread = p.head_dim / 32;
// Q: stride-based [batch, q_head, q_len=1, head_dim]
float q_reg[8]; float q_reg[8];
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
+ lane * hd_per_thread * p.q_stride_d; + lane * hd_per_thread * p.q_stride_d;
@@ -46,7 +38,8 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
int this_chunk = min(PDC_CHUNK, p.kv_len - chunk_start); int this_chunk = min(PDC_CHUNK, p.kv_len - chunk_start);
int total = this_chunk * p.head_dim; int total = this_chunk * p.head_dim;
for (int i = threadIdx.y * 32 + lane; i < total; i += blockDim.x * blockDim.y) { for (int i = threadIdx.y * 32 + lane; i < total;
i += blockDim.x * blockDim.y) {
int s = i / p.head_dim; int s = i / p.head_dim;
int d_dim = i % p.head_dim; int d_dim = i % p.head_dim;
int pos = chunk_start + s; int pos = chunk_start + s;
@@ -69,14 +62,19 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
float partial = 0.0f; float partial = 0.0f;
#pragma unroll #pragma unroll
for (int i = 0; i < hd_per_thread; i++) for (int i = 0; i < hd_per_thread; i++)
partial += q_reg[i] * __bfloat162float(k_smem[s * p.head_dim + lane * hd_per_thread + i]); partial += q_reg[i] * __bfloat162float(
partial = paged_warp_reduce_sum(partial) * p.scale; k_smem[s * p.head_dim + lane * hd_per_thread + i]);
partial = warp_reduce_sum(partial) * p.scale;
int kv_idx = chunk_start + s; int kv_idx = chunk_start + s;
if (p.use_mask && p.mask && !p.mask[mask_base + kv_idx]) if constexpr (HasMask) {
if (!p.mask[mask_base + kv_idx])
partial = -FLT_MAX; partial = -FLT_MAX;
if (p.causal_offset >= 0 && kv_idx > p.causal_offset) }
if constexpr (IsCausal) {
if (kv_idx > p.causal_offset)
partial = -FLT_MAX; partial = -FLT_MAX;
}
float new_m = fmaxf(m, partial); float new_m = fmaxf(m, partial);
float alpha = expf(m - new_m); float alpha = expf(m - new_m);
@@ -93,11 +91,12 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
+ (int64_t)kv_head * p.head_dim; + (int64_t)kv_head * p.head_dim;
#pragma unroll #pragma unroll
for (int i = 0; i < hd_per_thread; i++) for (int i = 0; i < hd_per_thread; i++)
acc_reg[i] = acc_reg[i] * alpha + __bfloat162float(p.v_cache[v_base + lane * hd_per_thread + i]) * beta; acc_reg[i] = fmaf(acc_reg[i], alpha,
__bfloat162float(p.v_cache[v_base + lane * hd_per_thread + i]) * beta);
} else { } else {
#pragma unroll #pragma unroll
for (int i = 0; i < hd_per_thread; i++) for (int i = 0; i < hd_per_thread; i++)
acc_reg[i] = acc_reg[i] * alpha + 0.0f * beta; acc_reg[i] = fmaf(acc_reg[i], alpha, 0.0f);
} }
m = new_m; m = new_m;
} }
@@ -105,7 +104,7 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
} }
size_t bh = (size_t)batch * p.q_head + q_head; size_t bh = (size_t)batch * p.q_head + q_head;
size_t slot = bh * p.num_splits + split; size_t slot = bh * MAX_SPLITS + split;
int d0 = lane * hd_per_thread; int d0 = lane * hd_per_thread;
#pragma unroll #pragma unroll
for (int i = 0; i < hd_per_thread; i++) for (int i = 0; i < hd_per_thread; i++)
@@ -124,7 +123,7 @@ __global__ void paged_attn_decode_combine_kernel(PagedAttentionParams<bf16> p) {
int batch = bh / p.q_head; int batch = bh / p.q_head;
int q_head = bh % p.q_head; int q_head = bh % p.q_head;
size_t split_base = (size_t)bh * p.num_splits; size_t split_base = (size_t)bh * MAX_SPLITS;
const float* mlp = p.ml_part + split_base * 2; const float* mlp = p.ml_part + split_base * 2;
const float* op = p.o_part + split_base * p.head_dim; const float* op = p.o_part + split_base * p.head_dim;
@@ -134,10 +133,10 @@ __global__ void paged_attn_decode_combine_kernel(PagedAttentionParams<bf16> p) {
if (mi <= -FLT_MAX) continue; if (mi <= -FLT_MAX) continue;
float li = mlp[s * 2 + 1]; float li = mlp[s * 2 + 1];
float nm = fmaxf(m, mi); float nm = fmaxf(m, mi);
float corr = __expf(m - nm); float corr = expf(m - nm);
float e = __expf(mi - nm); float e = expf(mi - nm);
acc = acc * corr + op[s * p.head_dim + d] * e; acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
l = l * corr + li * e; l = fmaf(l, corr, li * e);
m = nm; m = nm;
} }
+51 -60
View File
@@ -3,153 +3,144 @@
#include <cuda_bf16.h> #include <cuda_bf16.h>
#include "attn_common.h" #include "attn_common.h"
#include "attn_mma_utils.cuh" #include "attn_mma_utils.cuh"
#include "attn_warp_utils.cuh"
using bf16 = __nv_bfloat16;
// Paged split-KV tensor-core decode via GQA head-packing. // Paged split-KV tensor-core decode via GQA head-packing.
// Identical algorithm to attn_decode_split_kv_mma_kernel but reads K/V // Reads K/V directly from the page pool through a page table — one tile
// directly from the page pool through a page table, eliminating the gather // (BC=32) fits within a single page (page_size >= 32), so the page-table
// copy. Each tile (BC=32) fits within a single page (page_size >= 32), so // lookup happens once per tile for cp.async.
// the page-table lookup happens once per tile for cp.async. //
// IsCausal and HasMask are compile-time bools.
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1> template <typename Traits, bool IsCausal, bool HasMask>
__global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16> p) { __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16> p) {
constexpr int KD = HEAD_DIM / 16;
constexpr int NC8 = BC / 8;
constexpr int KT2 = BC / 16;
constexpr int DN8 = HEAD_DIM / 8;
constexpr int LD = HEAD_DIM;
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
constexpr int VEC = 8;
constexpr int TOTAL = BC * HEAD_DIM;
const int lane = threadIdx.x; const int lane = threadIdx.x;
const int gid = lane >> 2; const int gid = lane >> 2;
const int tid4 = lane & 3; const int tid4 = lane & 3;
const int kv_head = blockIdx.x; const int pass = blockIdx.x / p.kv_head;
const int kv_head = blockIdx.x % p.kv_head;
const int batch = blockIdx.y; const int batch = blockIdx.y;
const int split = blockIdx.z; const int split = blockIdx.z;
const int G = p.q_head / p.kv_head;
const int q_head0 = kv_head * G;
__shared__ __align__(16) bf16 sK[STAGES * BC * LD]; constexpr int MAX_G = 16;
__shared__ __align__(16) bf16 sV[STAGES * BC * LD]; const int G_total = p.q_head / p.kv_head;
const int g_begin = pass * MAX_G;
const int G = min(MAX_G, G_total - g_begin);
const int q_head0 = kv_head * G_total + g_begin;
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
// ---- Load Q directly from global into mma A-operand registers ----
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h; const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
const int qra = gid; const int qra = gid;
const int qrb = gid + 8; const int qrb = gid + 8;
const bool va = qra < G, vb = qrb < G; const bool va = qra < G, vb = qrb < G;
unsigned Qa[KD][4]; unsigned Qa[Traits::KD][4];
load_q_mma_frags<KD>(p.q + q_base, p.q_stride_h, p.q_stride_d, load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
qra, qrb, va, vb, tid4, Qa); qra, qrb, va, vb, tid4, Qa);
float Oacc[DN8][4]; float Oacc[Traits::DN8][4];
#pragma unroll #pragma unroll
for (int j = 0; j < DN8; j++) for (int j = 0; j < Traits::DN8; j++)
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f; Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f; float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
const int tiles_total = (p.kv_len + BC - 1) / BC; const int tiles_total = (p.kv_len + Traits::BC - 1) / Traits::BC;
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits; const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
const int ti_begin = split * tiles_per_split; const int ti_begin = split * tiles_per_split;
const int ti_end = min(tiles_total, ti_begin + tiles_per_split); const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
const int has_mask = p.use_mask && p.mask;
// Paged strides (constant for the block) const int64_t page_stride = (int64_t)p.page_size * p.kv_head * Traits::HEAD_DIM;
const int64_t page_stride = (int64_t)p.page_size * p.kv_head * HEAD_DIM; const int64_t pos_stride = (int64_t)p.kv_head * Traits::HEAD_DIM;
const int64_t pos_stride = (int64_t)p.kv_head * HEAD_DIM; const int64_t head_off = (int64_t)kv_head * Traits::HEAD_DIM;
const int64_t head_off = (int64_t)kv_head * HEAD_DIM;
// ---- Load tile lambda: predicated cp.async, paged addressing ---- // ---- Load tile lambda: paged addressing ----
auto load_tile = [&](int ti, int buf) { auto load_tile = [&](int ti, int buf) {
int kv0 = ti * BC; int kv0 = ti * Traits::BC;
bf16* dK = sK + buf * BC * LD; bf16* dK = sK + buf * Traits::BC * Traits::LD;
bf16* dV = sV + buf * BC * LD; bf16* dV = sV + buf * Traits::BC * Traits::LD;
int logical_page = kv0 / p.page_size; int logical_page = kv0 / p.page_size;
int phys_page = p.page_table[batch * p.max_pages + logical_page]; int phys_page = p.page_table[batch * p.max_pages + logical_page];
bool page_valid = (phys_page >= 0); bool page_valid = (phys_page >= 0);
#pragma unroll #pragma unroll
for (int i = lane * VEC; i < TOTAL; i += 32 * VEC) { for (int i = lane * Traits::VEC; i < Traits::TOTAL;
int r = i / HEAD_DIM, d = i % HEAD_DIM; i += Traits::NUM_THREADS * Traits::VEC) {
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
int kc = kv0 + r; int kc = kv0 + r;
bool valid = (kc < p.kv_len) && page_valid; bool valid = (kc < p.kv_len) && page_valid;
int page_off = kc % p.page_size; int page_off = kc % p.page_size;
int64_t gmem_base = (int64_t)phys_page * page_stride int64_t gmem_base = (int64_t)phys_page * page_stride
+ (int64_t)page_off * pos_stride + (int64_t)page_off * pos_stride
+ head_off; + head_off;
int off = r * LD + swiz_col(d, r, SWIZ_MASK); 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(&dK[off], &p.k_cache[gmem_base + d], valid);
cp_async_16_pred(&dV[off], &p.v_cache[gmem_base + d], valid); cp_async_16_pred(&dV[off], &p.v_cache[gmem_base + d], valid);
} }
cp_async_commit(); cp_async_commit();
}; };
// ---- Prologue: issue first tile load ---- constexpr int BUF_MASK = (Traits::STAGES > 1) ? (Traits::STAGES - 1) : 0;
if (ti_begin < ti_end) { if (ti_begin < ti_end) {
load_tile(ti_begin, 0); load_tile(ti_begin, 0);
} }
for (int ti = ti_begin; ti < ti_end; ti++) { for (int ti = ti_begin; ti < ti_end; ti++) {
constexpr int BUF_MASK = (STAGES > 1) ? (STAGES - 1) : 0;
int buf = (ti - ti_begin) & BUF_MASK; int buf = (ti - ti_begin) & BUF_MASK;
cp_async_wait_group<0>(); cp_async_wait_group<0>();
__syncwarp(); __syncwarp();
if constexpr (STAGES > 1) { if constexpr (Traits::STAGES > 1) {
if (ti + 1 < ti_end) if (ti + 1 < ti_end)
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK); load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
} }
const bf16* bK = sK + buf * BC * LD; const bf16* bK = sK + buf * Traits::BC * Traits::LD;
const bf16* bV = sV + buf * BC * LD; const bf16* bV = sV + buf * Traits::BC * Traits::LD;
int kv0 = ti * BC; int kv0 = ti * Traits::BC;
float Sacc[NC8][4]; float Sacc[Traits::NC8][4];
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc); mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
#pragma unroll #pragma unroll
for (int n8 = 0; n8 < NC8; n8++) for (int n8 = 0; n8 < Traits::NC8; n8++)
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale, Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale; Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
// Decode: q_len=1, so qrow0=qrow1=0, mask_q_stride irrelevant int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
int maxc = (p.causal_offset >= 0) ? min(p.kv_len, p.causal_offset + 1) : p.kv_len; mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
mma_softmax_tile<NC8, DN8>(kv0, maxc, maxc,
0, 0, 0, 0,
p.mask_b_stride, 0, p.mask_b_stride, 0,
batch, batch,
p.mask, has_mask, p.mask,
Sacc, Oacc, m0, m1, l0, l1, lane); Sacc, Oacc, m0, m1, l0, l1, lane);
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc); mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
__syncwarp(); __syncwarp();
if constexpr (STAGES == 1) { if constexpr (Traits::STAGES == 1) {
if (ti + 1 < ti_end) if (ti + 1 < ti_end)
load_tile(ti + 1, 0); load_tile(ti + 1, 0);
} }
} }
// ---- write UN-normalised partials for this split ----
auto split_slot = [&](int h) -> size_t { auto split_slot = [&](int h) -> size_t {
size_t bh = (size_t)batch * p.q_head + h; size_t bh = (size_t)batch * p.q_head + h;
return bh * p.num_splits + split; return bh * MAX_SPLITS + split;
}; };
#pragma unroll #pragma unroll
for (int dn8 = 0; dn8 < DN8; dn8++) { for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
int d = dn8 * 8 + 2 * tid4; int d = dn8 * 8 + 2 * tid4;
int r0 = gid, r1 = gid + 8; int r0 = gid, r1 = gid + 8;
if (r0 < G) { if (r0 < G) {
int h = q_head0 + r0; int h = q_head0 + r0;
float* op = p.o_part + split_slot(h) * HEAD_DIM; float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
op[d] = Oacc[dn8][0]; op[d] = Oacc[dn8][0];
op[d + 1] = Oacc[dn8][1]; op[d + 1] = Oacc[dn8][1];
} }
if (r1 < G) { if (r1 < G) {
int h = q_head0 + r1; int h = q_head0 + r1;
float* op = p.o_part + split_slot(h) * HEAD_DIM; float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
op[d] = Oacc[dn8][2]; op[d] = Oacc[dn8][2];
op[d + 1] = Oacc[dn8][3]; op[d + 1] = Oacc[dn8][3];
} }
+1 -30
View File
@@ -1,35 +1,6 @@
#include "attn_prefill_split_q.cuh" #include "attn_dispatchers.cuh"
#include "attn_entry_utils.cuh" #include "attn_entry_utils.cuh"
#ifndef ASTRAI_NO_MMA
#include "attn_prefill_split_q_mma.cuh"
#endif
template <int HEAD_DIM>
static void dispatch_prefill(AttentionParams<bf16>& p) {
#ifndef ASTRAI_NO_MMA
constexpr int WARPS = 4, BR = 16;
// KV tile: bigger tiles amortize the per-tile cp.async wait + barrier +
// loop overhead over more tensor-core work (this kernel is latency-bound,
// not compute/bandwidth-bound), so BC=32 wins ~6-8% over BC=16 for
// D<=128. D=256 stays at 16: BC=32 double-buffered would need 64KB smem,
// over the 48KB static cap.
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
dim3 grid((p.q_len + BR * WARPS - 1) / (BR * WARPS), p.q_head, p.batch);
dim3 block(WARPS * 32, 1, 1);
// Static shared memory — double-buffered K/V only (no sQ: Q goes direct
// to registers). 2*BC*LD bf16 each for sK and sV → 4*BC*HEAD_DIM*2 bytes.
// Occupancy is smem-capped: D=64→3 blocks/SM (16KB), D=128→1 (32KB),
// D=256→1 (32KB, BC=16).
attn_prefill_split_q_mma_kernel<HEAD_DIM, WARPS, BC><<<grid, block>>>(p);
#else
constexpr int G = 8, ROWS = 32, P_BC = 32;
dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
dim3 block(G, ROWS, 1);
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC><<<grid, block>>>(p);
#endif
}
torch::Tensor attn_prefill( torch::Tensor attn_prefill(
torch::Tensor q, torch::Tensor q,
torch::Tensor k, torch::Tensor k,
+14 -17
View File
@@ -6,12 +6,9 @@
using bf16 = __nv_bfloat16; using bf16 = __nv_bfloat16;
// v9: group-split register blocking. G threads cooperate on one query row, // v9: group-split register blocking. G threads cooperate on one query row,
// each owning HEAD_DIM/G dims of qreg[]/acc[]. Small per-thread footprint keeps // each owning HEAD_DIM/G dims of qreg[]/acc[]. IsCausal and HasMask are
// occupancy high; the S dot product is reduced across the G-lane group with a // compile-time bools — the compiler eliminates dead branches.
// short shuffle chain (log2(G) shuffles) instead of a full 32-lane warp reduce. // Templated on <HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask>.
// Online (per-kv) softmax — cheap because acc[] is only HEAD_DIM/G long.
// Templated on <HEAD_DIM, G, ROWS, P_BC>. Block = (G, ROWS). G power-of-two,
// G*ROWS a multiple of 32 with groups warp-aligned.
template <int G> template <int G>
__device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) { __device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
@@ -21,8 +18,7 @@ __device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
return v; return v;
} }
// load 8 contiguous bf16 from (16-byte aligned) smem as one float4, unpack to // load 8 contiguous bf16 from (16-byte aligned) smem as one float4
// 8 floats — cuts shared-load instructions 8x vs scalar bf16 loads.
__device__ __forceinline__ void ld8(const bf16* p, float* o) { __device__ __forceinline__ void ld8(const bf16* p, float* o) {
float4 raw = *reinterpret_cast<const float4*>(p); float4 raw = *reinterpret_cast<const float4*>(p);
const __nv_bfloat162* h = reinterpret_cast<const __nv_bfloat162*>(&raw); const __nv_bfloat162* h = reinterpret_cast<const __nv_bfloat162*>(&raw);
@@ -34,7 +30,7 @@ __device__ __forceinline__ void ld8(const bf16* p, float* o) {
} }
} }
template <int HEAD_DIM, int G, int ROWS, int P_BC> template <int HEAD_DIM, int G, int ROWS, int P_BC, bool IsCausal, bool HasMask>
__global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) { __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
constexpr int DPT = HEAD_DIM / G; constexpr int DPT = HEAD_DIM / G;
@@ -57,7 +53,7 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
+ q_row * p.q_stride_l + gpos * DPT * p.q_stride_d; + q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
#pragma unroll #pragma unroll
for (int i = 0; i < DPT; i++) for (int i = 0; i < DPT; i++)
qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]) * p.scale; qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
} }
float m = -FLT_MAX, l = 0.0f; float m = -FLT_MAX, l = 0.0f;
@@ -73,8 +69,6 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
int tt = G * ROWS; int tt = G * ROWS;
int lid = row * G + gpos; int lid = row * G + gpos;
// per-group shuffle mask: only the G lanes of this row's group participate,
// so causal masking (differing loop bounds across rows in a warp) is safe.
int lane_in_warp = lid & 31; int lane_in_warp = lid & 31;
unsigned gmask = (G == 32) ? 0xFFFFFFFFu unsigned gmask = (G == 32) ? 0xFFFFFFFFu
: (((1u << G) - 1u) << (lane_in_warp & ~(G - 1))); : (((1u << G) - 1u) << (lane_in_warp & ~(G - 1)));
@@ -95,13 +89,15 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
__syncthreads(); __syncthreads();
int lim = tlen; int lim = tlen;
if (p.causal_offset >= 0 && q_row < p.q_len) { if constexpr (IsCausal) {
if (q_row < p.q_len) {
int ep = q_row + p.causal_offset + 1; int ep = q_row + p.causal_offset + 1;
if (kv0 >= ep) if (kv0 >= ep)
lim = 0; lim = 0;
else if (kv0 + tlen > ep) else if (kv0 + tlen > ep)
lim = ep - kv0; lim = ep - kv0;
} }
}
int mask_row_base = mask_batch_base + q_row * p.mask_q_stride; int mask_row_base = mask_batch_base + q_row * p.mask_q_stride;
for (int s = 0; s < lim; s++) { for (int s = 0; s < lim; s++) {
@@ -115,11 +111,13 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
for (int j = 0; j < 8; j++) for (int j = 0; j < 8; j++)
part = fmaf(qreg[i + j], k8[j], part); part = fmaf(qreg[i + j], k8[j], part);
} }
float dot = group_reduce_sum<G>(part, gmask); float dot = group_reduce_sum<G>(part, gmask) * p.scale;
int kv_idx = kv0 + s; int kv_idx = kv0 + s;
if (p.use_mask && p.mask && !p.mask[mask_row_base + kv_idx]) if constexpr (HasMask) {
if (!p.mask[mask_row_base + kv_idx])
dot = -FLT_MAX; dot = -FLT_MAX;
}
float nm = fmaxf(m, dot); float nm = fmaxf(m, dot);
float al = __expf(m - nm); float al = __expf(m - nm);
@@ -141,10 +139,9 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
} }
if (q_row < p.q_len) { if (q_row < p.q_len) {
// O: stride-based write
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h int o_off = batch * p.q_stride_b + q_head * p.q_stride_h
+ q_row * p.q_stride_l + gpos * DPT * p.q_stride_d; + q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
float rl = (l > 1e-10f) ? (1.0f / l) : 0.0f; float rl = (l > 1e-20f) ? (1.0f / l) : 0.0f;
#pragma unroll #pragma unroll
for (int i = 0; i < DPT; i++) for (int i = 0; i < DPT; i++)
p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * rl); p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * rl);
+53 -103
View File
@@ -4,121 +4,76 @@
#include "attn_common.h" #include "attn_common.h"
#include "attn_mma_utils.cuh" #include "attn_mma_utils.cuh"
using bf16 = __nv_bfloat16;
// Tensor-core prefill flash attention (raw mma.sync PTX). // Tensor-core prefill flash attention (raw mma.sync PTX).
// One warp owns BR=16 query rows. S = Q@K^T and O = P@V run on bf16 tensor // One warp owns BR=16 query rows. S = Q@K^T and O = P@V run on bf16 tensor
// cores via mma.sync.m16n8k16 (f32 accumulate). Q fragments are loaded once // cores via mma.sync.m16n8k16 (f32 accumulate).
// straight from global into the mma A-operand layout (no smem staging) and
// kept resident in registers across the tile loop. S, O, and the online-softmax
// stats (m, l) also live in registers.
// Shared memory is statically sized via template parameters — no dynamic
// allocation. The mma fragment layout is used directly: the S accumulator
// (f32) maps element-for-element onto the P matrix_a (bf16) operand, so
// softmax needs no shuffle repack; row reductions fold across the 4-lane
// thread group. Templated on <HEAD_DIM, WARPS, BC> with BC a multiple of 16.
// //
// Software pipeline: K/V are double-buffered and loaded via cp.async one tile // IsCausal and HasMask are compile-time bools — the compiler eliminates all
// ahead, so the next tile streams from global memory while the current tile's // dead branches in the inner compute loop (FA2-style).
// tensor-core math runs — hiding load latency (long_scoreboard). A single
// __syncthreads per tile both publishes the freshly loaded tile cross-warp and
// (because it runs before the next prefetch) guards the buffer being refilled,
// so no second barrier is needed. Predicated cp.async (cp_async_16_pred)
// zero-fills rows past kv_len, unifying full and partial tiles on one path.
// BC=32 (D<=128) amortizes the per-tile wait+barrier+loop overhead over more
// tensor-core work — this kernel is latency-bound (low occupancy from high
// register pressure), so fewer, larger tiles beat many tiny ones.
// //
// Optimizations: load Q fragments directly from global in mma A-operand layout // Traits = KernelTraits<HEAD_DIM, BC, WARPS=4, STAGES=2>.
// (no sQ staging, no prologue barriers); post-multiply scale in float after template <typename Traits, bool IsCausal, bool HasMask>
// S=Q@K^T to avoid bf16 precision loss; packed bf16x2 output stores;
// causal tile skipping (block-level prefetch bound + warp-level compute skip);
// XOR swizzle (swiz_col) → eliminates ldmatrix bank conflicts without LD
// padding (LD=HEAD_DIM).
template <int HEAD_DIM, int WARPS, int BC>
__global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) { __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
constexpr int BR = 16;
constexpr int KD = HEAD_DIM / 16; // Q/K k-tiles
constexpr int NC8 = BC / 8; // S n-tiles (N=8 each)
constexpr int KT2 = BC / 16; // P k-tiles (K=16 each)
constexpr int DN8 = HEAD_DIM / 8; // O n-tiles (N=8 each)
constexpr int LD = HEAD_DIM; // XOR swizzle (swiz_col) handles bank conflicts
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1); // chunk bits, stay within LD
const int warp = threadIdx.x / 32; const int warp = threadIdx.x / 32;
const int lane = threadIdx.x % 32; const int lane = threadIdx.x % 32;
const int gid = lane >> 2; // 0..7 → rows gid, gid+8 const int gid = lane >> 2; // 0..7
const int tid4 = lane & 3; // 0..3 const int tid4 = lane & 3; // 0..3
const int nthreads = WARPS * 32;
const int q_head = blockIdx.y; const int q_head = blockIdx.y;
const int batch = blockIdx.z; const int batch = blockIdx.z;
const int kv_head = q_head / (p.q_head / p.kv_head); const int kv_head = q_head / (p.q_head / p.kv_head);
const int qrow0 = (blockIdx.x * WARPS + warp) * BR; const int qrow0 = (blockIdx.x * Traits::WARPS + warp) * Traits::BR;
// ---- Static shared memory: double-buffered K/V ---- // Static shared memory: double-buffered K/V (no sQ — Q goes direct
// K/V are double-buffered (STAGES=2): the next tile's cp.async load runs // to registers in mma A-operand layout).
// while the current tile's tensor-core math executes, hiding global-load __shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
// latency (FA2-style software pipeline). No dynamic smem / carveout opt-in. __shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
constexpr int STAGES = 2;
__shared__ __align__(16) bf16 sK[STAGES * BC * LD];
__shared__ __align__(16) bf16 sV[STAGES * BC * LD];
// Load Q fragments straight from global into mma A-operand layout. // Load Q fragments straight from global into mma A-operand layout.
// stride_row = p.q_stride_l for prefill (multi-q rows across q_len).
// See attn_mma_utils.cuh for the shared template.
const int q_base = batch * p.q_stride_b + q_head * p.q_stride_h; const int q_base = batch * p.q_stride_b + q_head * p.q_stride_h;
const int qra = qrow0 + gid; const int qra = qrow0 + gid;
const int qrb = qrow0 + gid + 8; const int qrb = qrow0 + gid + 8;
const bool va = qra < p.q_len, vb = qrb < p.q_len; const bool va = qra < p.q_len, vb = qrb < p.q_len;
unsigned Qa[KD][4]; unsigned Qa[Traits::KD][4];
load_q_mma_frags<KD>(p.q + q_base, p.q_stride_l, p.q_stride_d, load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_l, p.q_stride_d,
qra, qrb, va, vb, tid4, Qa); qra, qrb, va, vb, tid4, Qa);
float Oacc[DN8][4]; float Oacc[Traits::DN8][4];
#pragma unroll #pragma unroll
for (int j = 0; j < DN8; j++) for (int j = 0; j < Traits::DN8; j++)
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f; Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f; float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
// KV: stride-based base // KV: stride-based base
const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h; const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
const int tiles = (p.kv_len + BC - 1) / BC; const int tiles = (p.kv_len + Traits::BC - 1) / Traits::BC;
const int qr0 = qrow0 + gid; // row for c0/c1 const int qr0 = qrow0 + gid;
const int qr1 = qrow0 + gid + 8; // row for c2/c3 const int qr1 = qrow0 + gid + 8;
// Causal tile-skip bounds (no-op when causal_offset < 0) // Causal tile-skip bounds (dead code when IsCausal == false)
const int use_skip = (p.causal_offset >= 0) ? 1 : 0; const int max_kv = qrow0 + Traits::BR - 1 + p.causal_offset;
const int max_kv = qrow0 + BR - 1 + p.causal_offset;
const int block_max_kv = const int block_max_kv =
blockIdx.x * WARPS * BR + WARPS * BR - 1 + p.causal_offset; blockIdx.x * Traits::WARPS * Traits::BR + Traits::WARPS * Traits::BR - 1
const int has_mask = p.use_mask && p.mask; + p.causal_offset;
// Last active tile: block-level causal bound (all warps in the block share
// the K/V load, so the prefetch range is the block max, not per-warp).
int t_end = tiles - 1; int t_end = tiles - 1;
if (use_skip) { if constexpr (IsCausal) {
int bt = block_max_kv / BC; int bt = block_max_kv / Traits::BC;
if (bt < t_end) t_end = bt; if (bt < t_end) t_end = bt;
} }
constexpr int VEC = 8; // bf16 per cp.async unit (16 bytes)
constexpr int TOTAL = BC * HEAD_DIM;
// ---- Load tile lambda: predicated cp.async ---- // ---- Load tile lambda: predicated cp.async ----
// Issue cp.async loads for tile `ti` into shared buffer `buf`. Predicated
// loads zero-fill rows past kv_len, so partial tiles need no scalar path.
auto load_tile = [&](int ti, int buf) { auto load_tile = [&](int ti, int buf) {
int kv0 = ti * BC; int kv0 = ti * Traits::BC;
bf16* dK = sK + buf * BC * LD; bf16* dK = sK + buf * Traits::BC * Traits::LD;
bf16* dV = sV + buf * BC * LD; bf16* dV = sV + buf * Traits::BC * Traits::LD;
#pragma unroll #pragma unroll
for (int i = threadIdx.x * VEC; i < TOTAL; i += nthreads * VEC) { for (int i = threadIdx.x * Traits::VEC; i < Traits::TOTAL;
int r = i / HEAD_DIM, d = i % HEAD_DIM; i += Traits::NUM_THREADS * Traits::VEC) {
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
int kc = kv0 + r; int kc = kv0 + r;
bool valid = kc < p.kv_len; bool valid = kc < p.kv_len;
int off = r * LD + swiz_col(d, r, SWIZ_MASK); int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d; int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d;
cp_async_16_pred(&dK[off], &p.k[g_off], valid); cp_async_16_pred(&dK[off], &p.k[g_off], valid);
cp_async_16_pred(&dV[off], &p.v[g_off], valid); cp_async_16_pred(&dV[off], &p.v[g_off], valid);
@@ -132,65 +87,60 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
for (int ti = 0; ti <= t_end; ti++) { for (int ti = 0; ti <= t_end; ti++) {
int buf = ti & 1; int buf = ti & 1;
// Wait for the current tile's async copies, then a single barrier: it // Wait for current tile, then publish cross-warp + guard buffer reuse.
// both publishes this tile's data cross-warp AND guarantees the prior
// compute on the buffer we are about to refill has finished. Issuing
// the next tile's load *after* this barrier lets one barrier cover both
// hazards (vs two), while the load still overlaps this tile's math.
cp_async_wait_group<0>(); cp_async_wait_group<0>();
__syncthreads(); __syncthreads();
if (ti < t_end) load_tile(ti + 1, (ti + 1) & 1); if (ti < t_end) load_tile(ti + 1, (ti + 1) & 1);
const bf16* bK = sK + buf * BC * LD; const bf16* bK = sK + buf * Traits::BC * Traits::LD;
const bf16* bV = sV + buf * BC * LD; const bf16* bV = sV + buf * Traits::BC * Traits::LD;
int kv0 = ti * BC; int kv0 = ti * Traits::BC;
// Warp-level causal skip // Warp-level causal skip (dead branch eliminated when IsCausal == false)
if (!use_skip || kv0 <= max_kv) { if (!IsCausal || kv0 <= max_kv) {
// S = Q @ K^T + scale + online softmax + O += P @ V float Sacc[Traits::NC8][4];
float Sacc[NC8][4]; mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
// post-multiply scale in float (no bf16 precision loss from pre-scaling Q) // Post-multiply scale in float (no bf16 precision loss)
#pragma unroll #pragma unroll
for (int n8 = 0; n8 < NC8; n8++) for (int n8 = 0; n8 < Traits::NC8; n8++)
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale, Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale; Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
int maxc0 = (p.causal_offset >= 0) ? min(p.kv_len, qr0 + p.causal_offset + 1) int maxc0 = IsCausal ? min(p.kv_len, qr0 + p.causal_offset + 1)
: p.kv_len; : p.kv_len;
int maxc1 = (p.causal_offset >= 0) ? min(p.kv_len, qr1 + p.causal_offset + 1) int maxc1 = IsCausal ? min(p.kv_len, qr1 + p.causal_offset + 1)
: p.kv_len; : p.kv_len;
mma_softmax_tile<NC8, DN8>(kv0, maxc0, maxc1, mma_softmax_tile<Traits, HasMask>(kv0, maxc0, maxc1,
qr0, qr1, qr0, qr1,
p.mask_b_stride, p.mask_q_stride, p.mask_b_stride, p.mask_q_stride,
batch, batch,
p.mask, has_mask, p.mask,
Sacc, Oacc, m0, m1, l0, l1, lane); Sacc, Oacc, m0, m1, l0, l1, lane);
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc); mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
} // if active (warp-level causal skip) }
} }
// ---- write output ---- (packed bf16x2 stores: one 32-bit STG per pair, // ---- write output: packed bf16x2 stores ----
// halves store count and removes the uncoalesced scalar-store penalty)
float rl0 = (l0 > 1e-20f) ? (1.0f / l0) : 0.0f; float rl0 = (l0 > 1e-20f) ? (1.0f / l0) : 0.0f;
float rl1 = (l1 > 1e-20f) ? (1.0f / l1) : 0.0f; float rl1 = (l1 > 1e-20f) ? (1.0f / l1) : 0.0f;
// O: stride-based write
const int o_base = batch * p.q_stride_b + q_head * p.q_stride_h; const int o_base = batch * p.q_stride_b + q_head * p.q_stride_h;
#pragma unroll #pragma unroll
for (int dn8 = 0; dn8 < DN8; dn8++) { for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
int d = dn8 * 8 + 2 * tid4; int d = dn8 * 8 + 2 * tid4;
if (qr0 < p.q_len) { if (qr0 < p.q_len) {
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][0] * rl0, __nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][0] * rl0,
Oacc[dn8][1] * rl0); Oacc[dn8][1] * rl0);
*reinterpret_cast<__nv_bfloat162*>(&p.o[o_base + qr0 * p.q_stride_l + d * p.q_stride_d]) = v; *reinterpret_cast<__nv_bfloat162*>(
&p.o[o_base + qr0 * p.q_stride_l + d * p.q_stride_d]) = v;
} }
if (qr1 < p.q_len) { if (qr1 < p.q_len) {
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1, __nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1,
Oacc[dn8][3] * rl1); Oacc[dn8][3] * rl1);
*reinterpret_cast<__nv_bfloat162*>(&p.o[o_base + qr1 * p.q_stride_l + d * p.q_stride_d]) = v; *reinterpret_cast<__nv_bfloat162*>(
&p.o[o_base + qr1 * p.q_stride_l + d * p.q_stride_d]) = v;
} }
} }
} }
+13
View File
@@ -0,0 +1,13 @@
#pragma once
#include <cuda_bf16.h>
using bf16 = __nv_bfloat16;
static constexpr int MAX_SPLITS = 32;
__device__ inline float warp_reduce_sum(float val) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1)
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
return val;
}
+61 -77
View File
@@ -1,73 +1,33 @@
/* /*
Pure-C test: Pure-C test — uses shared dispatcher.
nvcc -I csrc -arch=sm_89 -O3 \ nvcc -I csrc -arch=sm_89 -O3 \
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \ --use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
csrc/tests/attn_decode_test.cu -o test && ./test csrc/tests/attn_decode_test.cu -o test && ./test
*/ */
#include "test_utils.cuh" #include "test_utils.cuh"
#include "../kernels/attn_decode_split_kv.cuh" #include "../kernels/attn_dispatchers.cuh"
#ifndef ASTRAI_NO_MMA
#include "../kernels/attn_decode_split_kv_mma.cuh"
#endif
// Split-K scratch (torch-free): the production launcher allocates these from // Split-K scratch (torch-free)
// torch; here we pass pre-allocated device buffers so the bench loop doesn't
// pay a cudaMalloc per iteration. Size for the maximum split count (32).
struct DecodeScratch { struct DecodeScratch {
float* o_part = nullptr; float* o_part = nullptr;
float* ml_part = nullptr; float* ml_part = nullptr;
}; };
// Launch the production decode path (tensor-core head-packing MMA on sm_80+, static void setup_scratch(AttentionParams<bf16>& p, DecodeScratch& sc) {
// scalar fallback otherwise), mirroring dispatch_decode() in attn_decode.cu. int max_splits = 32;
#ifndef ASTRAI_NO_MMA cudaMalloc(&sc.o_part, (size_t)p.batch * p.q_head * max_splits * p.head_dim * sizeof(float));
static bool decode_use_mma(const AttentionParams<bf16>& p) { cudaMalloc(&sc.ml_part, (size_t)p.batch * p.q_head * max_splits * 2 * sizeof(float));
int G = p.q_head / p.kv_head;
return !p.use_mask && G > 1 && G <= 16;
} }
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1> static void free_scratch(DecodeScratch& sc) {
static void launch_mma_decode(AttentionParams<bf16>& p, DecodeScratch& sc) { cudaFree(sc.o_part); cudaFree(sc.ml_part);
int tiles_total = (p.kv_len + BC - 1) / BC;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
p.o_part = sc.o_part;
p.ml_part = sc.ml_part;
attn_decode_split_kv_mma_kernel<HEAD_DIM, BC, STAGES>
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
#endif
static void launch_scalar_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
int gs = p.q_head / p.kv_head;
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
p.o_part = sc.o_part;
p.ml_part = sc.ml_part;
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
attn_decode_split_kv_kernel<<<dim3(p.batch * p.kv_head, 1, p.num_splits), dim3(32, gs), smem>>>(p);
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
template <int HEAD_DIM>
static void dispatch_decode_t(AttentionParams<bf16>& p, DecodeScratch& sc) {
#ifndef ASTRAI_NO_MMA
if (decode_use_mma(p)) { launch_mma_decode<HEAD_DIM, 32>(p, sc); return; }
#endif
launch_scalar_decode(p, sc);
}
static void dispatch_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
dispatch_by_head_dim(p.head_dim, [&]<int D>() { dispatch_decode_t<D>(p, sc); });
} }
// Warmed-up, CUDA-event timed sweep over the production decode MMA path. // Warmed-up, CUDA-event timed sweep over the production decode MMA path.
static void bench() { static void bench() {
const int cfgs[][5] = { const int cfgs[][5] = {
{1, 32, 4, 512, 128}, // B, Hq, Hk, kv_len, D {1, 32, 4, 512, 128},
{1, 32, 4, 1024, 128}, {1, 32, 4, 1024, 128},
{1, 32, 4, 2048, 128}, {1, 32, 4, 2048, 128},
{1, 32, 4, 4096, 128}, {1, 32, 4, 4096, 128},
@@ -104,10 +64,10 @@ static void bench() {
p.q = dQ; p.k = dK; p.v = dV; p.mask = nullptr; p.o = dO; p.q = dQ; p.k = dK; p.v = dV; p.mask = nullptr; p.o = dO;
DecodeScratch sc; DecodeScratch sc;
cudaMalloc(&sc.o_part, (size_t)B*Hq*32*D*sizeof(float)); setup_scratch(p, sc);
cudaMalloc(&sc.ml_part, (size_t)B*Hq*32*2*sizeof(float)); p.o_part = sc.o_part; p.ml_part = sc.ml_part;
auto launch = [&]() { dispatch_decode(p, sc); }; auto launch = [&]() { dispatch_by_head_dim(D, [&]<int H>() { dispatch_decode<H>(p); }); };
double flops = 4.0 * B * Hq * (double)sl * D; double flops = 4.0 * B * Hq * (double)sl * D;
double bytes = 2.0 * (2.0 * nKV * sizeof(bf16)); double bytes = 2.0 * (2.0 * nKV * sizeof(bf16));
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops, bytes); BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops, bytes);
@@ -119,22 +79,14 @@ static void bench() {
print_bench_row(cfg, r); print_bench_row(cfg, r);
cudaFree(dQ); cudaFree(dK); cudaFree(dV); cudaFree(dO); cudaFree(dQ); cudaFree(dK); cudaFree(dV); cudaFree(dO);
cudaFree(sc.o_part); cudaFree(sc.ml_part); free_scratch(sc);
} }
} }
int main() { static int run_test(int B, int Hq, int Hk, int sl, int D, int causal) {
const int configs[][5] = { int gs = Hq / Hk;
{1, 2, 1, 64, 32}, // B,Hq,Hk,seq_len,D printf("=== B=%d Hq=%d Hk=%d seq=%d D=%d gs=%d causal=%d ===\n",
{1, 32, 4, 512, 128}, B,Hq,Hk,sl,D,gs,causal);
{1, 32, 4, 1024, 128},
};
int n_cfgs = sizeof(configs) / sizeof(configs[0]);
for (int ci = 0; ci < n_cfgs; ci++) {
int B = configs[ci][0], Hq = configs[ci][1], Hk = configs[ci][2];
int sl = configs[ci][3], D = configs[ci][4], gs = Hq / Hk;
printf("=== B=%d Hq=%d Hk=%d seq=%d D=%d gs=%d ===\n", B,Hq,Hk,sl,D,gs);
size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D; size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D;
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV]; float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
@@ -161,18 +113,17 @@ int main() {
AttentionParams<bf16> p; AttentionParams<bf16> p;
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D; p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D;
p.use_mask=0; p.causal_offset=-1; p.use_mask=0; p.causal_offset=causal?0:-1;
p.scale=1.0f/sqrtf((float)D); p.scale=1.0f/sqrtf((float)D);
set_default_strides(p); set_default_strides(p);
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO; p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
// Split-K scratch (max 32 splits), sized for the production MMA path.
DecodeScratch sc; DecodeScratch sc;
cudaMalloc(&sc.o_part, (size_t)B*Hq*32*D*sizeof(float)); setup_scratch(p, sc);
cudaMalloc(&sc.ml_part, (size_t)B*Hq*32*2*sizeof(float)); p.o_part = sc.o_part; p.ml_part = sc.ml_part;
double t0=now_ms(); double t0=now_ms();
dispatch_decode(p, sc); dispatch_by_head_dim(D, [&]<int H>() { dispatch_decode<H>(p); });
cudaDeviceSynchronize(); cudaDeviceSynchronize();
double kms=now_ms()-t0; double kms=now_ms()-t0;
cudaError_t err=cudaGetLastError(); cudaError_t err=cudaGetLastError();
@@ -182,18 +133,51 @@ int main() {
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost); cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
float* ref=new float[nQ]; float* ref=new float[nQ];
cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, -1); cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, causal ? 0 : -1);
float max_err=0; float max_abs_err=0, max_rel_err=0;
for (size_t i=0;i<nQ;i++){ for (size_t i=0;i<nQ;i++){
float d=fabsf(bf2f(hOut[i])-ref[i]); float err=fabsf(bf2f(hOut[i])-ref[i]);
if(d>max_err) max_err=d; 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;
} }
printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err); 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); cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask);
cudaFree(sc.o_part);cudaFree(sc.ml_part); free_scratch(sc);
delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp; 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"); printf("All tests passed!\n");
bench(); bench();
+53 -77
View File
@@ -5,12 +5,8 @@
#include <cstring> #include <cstring>
#include "test_utils.cuh" #include "test_utils.cuh"
#include "../kernels/attn_paged_decode_split_kv.cuh" #include "../kernels/attn_dispatchers.cuh"
#ifndef ASTRAI_NO_MMA
#include "../kernels/attn_paged_decode_split_kv_mma.cuh"
#endif
// Copy contiguous K/V from page pool (reference gather)
static void gather_kv_cpu( static void gather_kv_cpu(
const bf16* h_k_pool, const bf16* h_v_pool, const bf16* h_k_pool, const bf16* h_v_pool,
const int64_t* h_pt, int B, int Hkv, int kv_len, const int64_t* h_pt, int B, int Hkv, int kv_len,
@@ -28,7 +24,8 @@ static void gather_kv_cpu(
size_t src_base = (size_t)phys * page_stride size_t src_base = (size_t)phys * page_stride
+ (size_t)pg_off * Hkv * head_dim + (size_t)pg_off * Hkv * head_dim
+ h * head_dim; + h * head_dim;
size_t dst_base = ((size_t)b * Hkv + h) * kv_len * head_dim + (size_t)pos * 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_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)); memcpy(h_v + dst_base, h_v_pool + src_base, head_dim * sizeof(bf16));
} }
@@ -37,54 +34,29 @@ static void gather_kv_cpu(
} }
template <int HEAD_DIM> template <int HEAD_DIM>
static void launch_paged_decode(PagedAttentionParams<bf16, float>& p) { static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int causal, int seed) {
#ifndef ASTRAI_NO_MMA printf("B=%d Hq=%d Hkv=%d kv_len=%d page_sz=%d head_dim=%d causal=%d ... ",
int G_check = p.q_head / p.kv_head; B, Hq, Hkv, kv_len, page_size, HEAD_DIM, causal);
bool use_mma = !p.use_mask && G_check >= 1 && G_check <= 16 && p.page_size >= 32;
if (use_mma) {
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
int tiles_total = (p.kv_len + 32 - 1) / 32;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
paged_attn_decode_split_kv_mma_kernel<HEAD_DIM, 32, STAGES>
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
} else
#endif
{
int group_sz = p.q_head / p.kv_head;
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
paged_attn_decode_split_kv_kernel<<<
dim3(p.batch * p.kv_head, 1, p.num_splits),
dim3(32, group_sz), smem>>>(p);
}
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
template <int HEAD_DIM>
static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed) {
printf("B=%d Hq=%d Hkv=%d kv_len=%d page_sz=%d head_dim=%d ... ", B, Hq, Hkv, kv_len, page_size, HEAD_DIM);
fflush(stdout); fflush(stdout);
int max_pages = (kv_len + page_size - 1) / page_size; int max_pages = (kv_len + page_size - 1) / page_size;
int n_phys_pages = B * max_pages; 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_q = (size_t)B * Hq * 1 * HEAD_DIM * sizeof(bf16);
size_t sz_o = sz_q; 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_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_pt = (size_t)B * max_pages * sizeof(int64_t);
int max_splits = 32;
size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float); size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float);
size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float); size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float);
bf16 *d_q, *d_o_paged, *d_o_ref; bf16 *d_q, *d_o_paged;
bf16 *d_k_pool, *d_v_pool; bf16 *d_k_pool, *d_v_pool;
int64_t* d_pt; int64_t* d_pt;
float *d_op, *d_ml; float *d_op, *d_ml;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_q, sz_q);
cudaMalloc(&d_o_paged, sz_o); cudaMalloc(&d_o_paged, sz_o);
cudaMalloc(&d_o_ref, sz_o);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_k_pool, sz_kv);
cudaMalloc(&d_v_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
cudaMalloc(&d_pt, sz_pt); cudaMalloc(&d_pt, sz_pt);
@@ -107,7 +79,8 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed)
for (int h = 0; h < Hkv; h++) { for (int h = 0; h < Hkv; h++) {
for (int d = 0; d < HEAD_DIM; d++) { for (int d = 0; d < HEAD_DIM; d++) {
float v = sinf((float)(pg * 7919 + off * 1049 + h * 331 + 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; size_t idx = (size_t)pg * ps + (size_t)off * Hkv * HEAD_DIM
+ h * HEAD_DIM + d;
h_k_pool[idx] = __float2bfloat16(v); h_k_pool[idx] = __float2bfloat16(v);
h_v_pool[idx] = __float2bfloat16(v * 0.3f); h_v_pool[idx] = __float2bfloat16(v * 0.3f);
} }
@@ -138,22 +111,22 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed)
} }
float* h_o_ref = (float*)calloc(B * Hq * HEAD_DIM, sizeof(float)); float* h_o_ref = (float*)calloc(B * Hq * HEAD_DIM, sizeof(float));
cpu_attention_ref(h_q_f, h_k_f, h_v_f, nullptr, h_o_ref, B, Hq, Hkv, 1, kv_len, HEAD_DIM, -1); 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);
float scale_val = 1.0f / sqrtf((float)HEAD_DIM); PagedAttentionParams<bf16> p;
PagedAttentionParams<bf16, float> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.q_len = 1; 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.kv_len = kv_len; p.head_dim = HEAD_DIM;
p.use_mask = 0; p.causal_offset = -1; p.use_mask = 0; p.causal_offset = causal ? 0 : -1;
set_default_paged_strides(p); set_default_paged_strides(p);
p.num_splits = 1; p.scale = scale_val; p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.page_size = page_size; p.max_pages = max_pages; p.page_size = page_size; p.max_pages = max_pages;
p.page_table = d_pt; p.page_table = d_pt;
p.k_cache = d_k_pool; p.v_cache = d_v_pool; 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.q = d_q; p.mask = nullptr; p.o = d_o_paged;
p.o_part = d_op; p.ml_part = d_ml; p.o_part = d_op; p.ml_part = d_ml;
launch_paged_decode<HEAD_DIM>(p); dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(p); });
cudaDeviceSynchronize(); cudaDeviceSynchronize();
bf16* h_o_bf16 = (bf16*)malloc(sz_o); bf16* h_o_bf16 = (bf16*)malloc(sz_o);
@@ -162,23 +135,30 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed)
for (int i = 0; i < B * Hq * HEAD_DIM; i++) for (int i = 0; i < B * Hq * HEAD_DIM; i++)
h_o_paged[i] = __bfloat162float(h_o_bf16[i]); h_o_paged[i] = __bfloat162float(h_o_bf16[i]);
float max_err = 0.0f; float max_abs_err = 0.0f, max_rel_err = 0.0f;
int bad_idx = -1; int bad_idx = -1;
for (int i = 0; i < B * Hq * HEAD_DIM; i++) { for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
float e = fabsf(h_o_paged[i] - h_o_ref[i]); float e = fabsf(h_o_paged[i] - h_o_ref[i]);
if (e > max_err) { max_err = e; bad_idx = 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;
} }
bool pass = max_err < 0.02f; 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) { if (pass) {
printf("PASS (max_abs_err=%.4e)\n", max_err); printf("PASS (max_abs_err=%.4e max_rel_err=%.4e)\n", max_abs_err, max_rel_err);
} else { } else {
int b = bad_idx / (Hq * HEAD_DIM); int b = bad_idx / (Hq * HEAD_DIM);
int h = (bad_idx / HEAD_DIM) % Hq; int h = (bad_idx / HEAD_DIM) % Hq;
int d = bad_idx % HEAD_DIM; int d = bad_idx % HEAD_DIM;
printf("FAIL (max_abs_err=%.4e at [%d,%d,%d]: ref=%.4f got=%.4f)\n", printf("FAIL (max_abs_err=%.4e max_rel_err=%.4e at [%d,%d,%d]: ref=%.4f got=%.4f)\n",
max_err, b, h, d, h_o_ref[bad_idx], h_o_paged[bad_idx]); max_abs_err, max_rel_err, b, h, d, h_o_ref[bad_idx], h_o_paged[bad_idx]);
printf(" ref[0..7]:"); printf(" ref[0..7]:");
for (int i = 0; i < 8 && i < HEAD_DIM; i++) for (int i = 0; i < 8 && i < HEAD_DIM; i++)
printf(" %.4f", h_o_ref[i]); printf(" %.4f", h_o_ref[i]);
@@ -192,7 +172,7 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed)
free(h_k_cont); free(h_v_cont); free(h_k_cont); free(h_v_cont);
free(h_q_f); free(h_k_f); free(h_v_f); free(h_q_f); free(h_k_f); free(h_v_f);
free(h_o_ref); free(h_o_bf16); free(h_o_paged); free(h_o_ref); free(h_o_bf16); free(h_o_paged);
cudaFree(d_q); cudaFree(d_o_paged); cudaFree(d_o_ref); cudaFree(d_q); cudaFree(d_o_paged);
cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_pt); cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_pt);
cudaFree(d_op); cudaFree(d_ml); cudaFree(d_op); cudaFree(d_ml);
@@ -201,48 +181,43 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed)
struct TestCase { struct TestCase {
int head_dim; int head_dim;
int B, Hq, Hkv, kv_len, page_size, seed; int B, Hq, Hkv, kv_len, page_size, causal, seed;
}; };
static const TestCase TESTS[] = { static const TestCase TESTS[] = {
{128, 1, 1, 1, 8, 128, 1}, {128, 1, 1, 1, 8, 128, 0, 1},
{128, 1, 4, 4, 128, 128, 2}, {128, 1, 4, 4, 128, 128, 0, 2},
{128, 2, 4, 4, 256, 128, 3}, {128, 2, 4, 4, 256, 128, 0, 3},
{128, 1, 4, 1, 64, 64, 4}, {128, 1, 4, 1, 64, 64, 0, 4},
{128, 1, 8, 2, 64, 128, 5}, {128, 1, 8, 2, 64, 128, 0, 5},
{128, 2, 16, 4, 128, 128, 6}, {128, 2, 16, 4, 128, 128, 0, 6},
{64, 1, 4, 2, 32, 128, 7}, {64, 1, 4, 2, 32, 128, 0, 7},
{256, 1, 2, 1, 16, 128, 8}, {256, 1, 2, 1, 16, 128, 0, 8},
{32, 1, 4, 2, 32, 64, 9}, {32, 1, 4, 2, 32, 64, 0, 9},
{128, 3, 8, 2, 256, 128, 10}, {128, 3, 8, 2, 256, 128, 0, 10},
{128, 2, 32, 8, 512, 128, 11}, {128, 2, 32, 8, 512, 128, 0, 11},
#ifndef ASTRAI_NO_MMA {128, 1, 16, 2, 256, 128, 0, 12},
{128, 1, 16, 2, 256, 128, 12}, {128, 2, 32, 4, 512, 128, 0, 13},
{128, 2, 32, 4, 512, 128, 13}, {128, 2, 8, 2, 128, 128, 1, 14}, // causal
#endif
}; };
static int dispatch_test(const TestCase& tc) { static int dispatch_test(const TestCase& tc) {
bool matched = false;
int r = 0; int r = 0;
dispatch_by_head_dim(tc.head_dim, [&]<int D>() { dispatch_by_head_dim(tc.head_dim, [&]<int D>() {
matched = true; r = run_test<D>(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.causal, tc.seed);
r = run_test<D>(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.seed);
}); });
return matched ? r : 1; return r;
} }
// Warmed-up, CUDA-event timed sweep over paged decode configs.
// Bytes = K + V read through page table (B*Hk*kv*D each), bf16.
template <int HEAD_DIM> template <int HEAD_DIM>
static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) { 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 max_pages = (kv_len + page_size - 1) / page_size;
int n_phys_pages = B * max_pages; 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_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_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_pt = (size_t)B * max_pages * sizeof(int64_t);
int max_splits = 32;
size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float); size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float);
size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float); size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float);
@@ -269,13 +244,12 @@ static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) {
cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice); cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice);
free(h_pt); free(h_pt);
float scale_val = 1.0f / sqrtf((float)HEAD_DIM); PagedAttentionParams<bf16> pa;
PagedAttentionParams<bf16, float> pa;
pa.batch = B; pa.q_head = Hq; pa.kv_head = Hkv; pa.q_len = 1; 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.kv_len = kv_len; pa.head_dim = HEAD_DIM;
pa.use_mask = 0; pa.causal_offset = -1; pa.use_mask = 0; pa.causal_offset = -1;
set_default_paged_strides(pa); set_default_paged_strides(pa);
pa.num_splits = 1; pa.scale = scale_val; pa.scale = 1.0f / sqrtf((float)HEAD_DIM);
pa.page_size = page_size; pa.max_pages = max_pages; pa.page_size = page_size; pa.max_pages = max_pages;
pa.page_table = d_pt; pa.page_table = d_pt;
pa.k_cache = d_k_pool; pa.v_cache = d_v_pool; pa.k_cache = d_k_pool; pa.v_cache = d_v_pool;
@@ -283,7 +257,9 @@ static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) {
pa.o_part = d_op; pa.ml_part = d_ml; pa.o_part = d_op; pa.ml_part = d_ml;
const int WARMUP = 10, ITERS = 100; const int WARMUP = 10, ITERS = 100;
auto launch = [&]() { launch_paged_decode<HEAD_DIM>(pa); }; 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; double flops = 4.0 * B * Hq * (double)kv_len * HEAD_DIM;
size_t nKV = (size_t)B * Hkv * kv_len * HEAD_DIM; size_t nKV = (size_t)B * Hkv * kv_len * HEAD_DIM;
double bytes = 2.0 * (2.0 * nKV * sizeof(bf16)); double bytes = 2.0 * (2.0 * nKV * sizeof(bf16));
+45 -54
View File
@@ -1,45 +1,14 @@
/* /*
Pure-C test: Pure-C test — uses shared dispatcher.
nvcc -I csrc -arch=sm_89 -O3 \ nvcc -I csrc -arch=sm_89 -O3 \
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \ --use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
csrc/tests/attn_prefill_test.cu -o test && ./test csrc/tests/attn_prefill_test.cu -o test && ./test
*/ */
#include "test_utils.cuh" #include "test_utils.cuh"
#include "../kernels/attn_prefill_split_q.cuh" #include "../kernels/attn_dispatchers.cuh"
#ifndef ASTRAI_NO_MMA
#include "../kernels/attn_prefill_split_q_mma.cuh"
#endif
// Launch the production prefill path (tensor-core MMA on sm_80+, else the
// scalar fallback), mirroring dispatch_prefill() in attn_prefill.cu.
template <int HEAD_DIM>
static void launch_prefill(AttentionParams<bf16>& p) {
#ifndef ASTRAI_NO_MMA
constexpr int WARPS = 4, BR = 16;
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
dim3 grid((p.q_len + BR * WARPS - 1) / (BR * WARPS), p.q_head, p.batch);
dim3 block(WARPS * 32, 1, 1);
attn_prefill_split_q_mma_kernel<HEAD_DIM, WARPS, BC><<<grid, block>>>(p);
#else
constexpr int G = 8, ROWS = 32, P_BC = 32;
dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
dim3 block(G, ROWS, 1);
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC><<<grid, block>>>(p);
#endif
}
static void dispatch_prefill(AttentionParams<bf16>& p) {
switch (p.head_dim) {
case 64: launch_prefill<64>(p); break;
case 128: launch_prefill<128>(p); break;
default: printf("bench: unsupported D=%d\n", p.head_dim);
}
}
// Warmed-up, CUDA-event timed throughput sweep over the production MMA path. // Warmed-up, CUDA-event timed throughput sweep over the production MMA path.
// Reports per-call latency and effective tensor-core TFLOP/s (2 matmuls:
// QK^T and P@V, each 2*B*Hq*ql*kl*D flops; halved for causal).
static void bench() { static void bench() {
const int cfgs[][7] = { const int cfgs[][7] = {
{1,32,4,512,512,128,0}, {1,32,4,512,512,128,0},
@@ -80,21 +49,21 @@ static void bench() {
p.scale=1.0f/sqrtf((float)D); p.scale=1.0f/sqrtf((float)D);
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO; p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
for (int i=0;i<WARMUP;i++) dispatch_prefill(p); auto launch = [&]() { dispatch_by_head_dim(D, [&]<int H>() { dispatch_prefill<H>(p); }); };
for (int i=0;i<WARMUP;i++) launch();
cudaDeviceSynchronize(); cudaDeviceSynchronize();
cudaError_t err=cudaGetLastError(); cudaError_t err=cudaGetLastError();
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return;} if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return;}
cudaEvent_t s,e; cudaEventCreate(&s); cudaEventCreate(&e); cudaEvent_t s,e; cudaEventCreate(&s); cudaEventCreate(&e);
cudaEventRecord(s); cudaEventRecord(s);
for (int i=0;i<ITERS;i++) dispatch_prefill(p); for (int i=0;i<ITERS;i++) launch();
cudaEventRecord(e); cudaEventSynchronize(e); cudaEventRecord(e); cudaEventSynchronize(e);
float ms=0; cudaEventElapsedTime(&ms,s,e); ms/=ITERS; float ms=0; cudaEventElapsedTime(&ms,s,e); ms/=ITERS;
double flops = 4.0*B*Hq*(double)ql*kl*D; double flops = 4.0*B*Hq*(double)ql*kl*D;
if (causal) flops *= 0.5; if (causal) flops *= 0.5;
double tflops = flops/(ms*1e-3)/1e12; double tflops = flops/(ms*1e-3)/1e12;
// HBM traffic: Q + O (B*Hq*ql*D each) + K + V (B*Hk*kl*D each), bf16.
double bytes = 2.0 * (2.0*nQ + 2.0*nKV); double bytes = 2.0 * (2.0*nQ + 2.0*nKV);
double gbps = bytes/(ms*1e-3)/1e9; double gbps = bytes/(ms*1e-3)/1e9;
@@ -110,19 +79,7 @@ static void bench() {
} }
} }
int main() { static int run_test(int B, int Hq, int Hk, int ql, int kl, int D, int causal) {
const int configs[][7] = {
{1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal
{1,32,4,512,512,128,0}, // standard
{1,32,4,128,256,128,0}, // medium
{1,4,2,256,256,128,1}, // causal
};
int n_configs = sizeof(configs) / sizeof(configs[0]);
for (int ci = 0; ci < n_configs; ci++) {
int B=configs[ci][0], Hq=configs[ci][1], Hk=configs[ci][2];
int ql=configs[ci][3], kl=configs[ci][4], D=configs[ci][5];
int causal=configs[ci][6];
printf("=== B=%d Hq=%d Hk=%d q=%d kv=%d D=%d causal=%d ===\n", printf("=== B=%d Hq=%d Hk=%d q=%d kv=%d D=%d causal=%d ===\n",
B,Hq,Hk,ql,kl,D,causal); B,Hq,Hk,ql,kl,D,causal);
@@ -150,7 +107,7 @@ int main() {
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO; p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
double t0=now_ms(); double t0=now_ms();
dispatch_prefill(p); dispatch_by_head_dim(D, [&]<int H>() { dispatch_prefill<H>(p); });
cudaDeviceSynchronize(); cudaDeviceSynchronize();
double kms=now_ms()-t0; double kms=now_ms()-t0;
cudaError_t err=cudaGetLastError(); cudaError_t err=cudaGetLastError();
@@ -162,15 +119,49 @@ int main() {
float* ref=new float[nQ]; float* ref=new float[nQ];
cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal ? 0 : -1); cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal ? 0 : -1);
float max_err=0; float max_abs_err=0, max_rel_err=0;
for (size_t i=0;i<nQ;i++) { for (size_t i=0;i<nQ;i++) {
float d=fabsf(bf2f(hOut[i])-ref[i]); float err=fabsf(bf2f(hOut[i])-ref[i]);
if(d>max_err) max_err=d; 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;
} }
printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err); 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(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp; 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"); printf("All tests passed!\n");
bench(); bench();
-10
View File
@@ -18,16 +18,6 @@ inline double now_ms() {
return duration_cast<milliseconds>(steady_clock::now().time_since_epoch()).count(); return duration_cast<milliseconds>(steady_clock::now().time_since_epoch()).count();
} }
inline int compute_num_splits(int base_blocks, int tiles_total) {
int sm_count = 0;
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
if (n > tiles_total) n = tiles_total;
if (n > 32) n = 32;
if (n < 1) n = 1;
return n;
}
#define CUDA_CHECK(call) \ #define CUDA_CHECK(call) \
do { \ do { \
cudaError_t _e = (call); \ cudaError_t _e = (call); \
+1
View File
@@ -50,3 +50,4 @@ quote-style = "double"
indent-style = "space" indent-style = "space"
skip-magic-trailing-comma = false skip-magic-trailing-comma = false
line-ending = "auto" line-ending = "auto"
exclude = ["*.md", "*.json", "*.yml", "*.yaml"]
+1 -1
View File
@@ -5,7 +5,7 @@ from huggingface_hub import snapshot_download
PROJECT_ROOT = Path(__file__).resolve().parents[2] PROJECT_ROOT = Path(__file__).resolve().parents[2]
DEFAULT_LOCAL_DIR = Path(PROJECT_ROOT, "params") DEFAULT_LOCAL_DIR = Path(PROJECT_ROOT, "params")
DEFAULT_REPO_ID = "ViperEk/KHAOSZ" DEFAULT_REPO_ID = "ViperEkura/AstrAI-V1-instruct"
if __name__ == "__main__": if __name__ == "__main__":
parser = argparse.ArgumentParser( parser = argparse.ArgumentParser(
+2 -4
View File
@@ -26,11 +26,9 @@ def batch_generate():
prompts = [ prompts = [
tokenizer.apply_chat_template( tokenizer.apply_chat_template(
[ [{"role": "user", "content": q}],
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": q},
],
tokenize=False, tokenize=False,
add_generation_prompt=True,
) )
for q in inputs for q in inputs
] ]
+9 -8
View File
@@ -58,8 +58,8 @@ def parse_args():
parser.add_argument( parser.add_argument(
"--system_prompt", "--system_prompt",
type=str, type=str,
default="You are a helpful assistant.", default="",
help="Optional system prompt", help="Optional system prompt (default: empty, model not SFT-trained on system role)",
) )
return parser.parse_args() return parser.parse_args()
@@ -73,18 +73,20 @@ def chat():
model.to(device="cuda", dtype=torch.bfloat16) model.to(device="cuda", dtype=torch.bfloat16)
engine = InferenceEngine(model=model, tokenizer=tokenizer) engine = InferenceEngine(model=model, tokenizer=tokenizer)
messages = [{"role": "system", "content": args.system_prompt}]
while True: while True:
query = input(">> ") query = input(">> ")
if query == "!exit": if query == "!exit":
break break
messages.append({"role": "user", "content": query}) msgs = []
if args.system_prompt:
msgs.append({"role": "system", "content": args.system_prompt})
msgs.append({"role": "user", "content": query})
prompt = tokenizer.apply_chat_template(
msgs, tokenize=False, add_generation_prompt=True
)
full_response = "" full_response = ""
prompt = tokenizer.apply_chat_template(messages, tokenize=False)
for token in engine.generate( for token in engine.generate(
prompt=prompt, prompt=prompt,
stream=True, stream=True,
@@ -99,7 +101,6 @@ def chat():
full_response += token full_response += token
print() print()
messages.append({"role": "assistant", "content": full_response.strip()})
if __name__ == "__main__": if __name__ == "__main__":
+3 -3
View File
@@ -143,7 +143,7 @@ def print_layer_grid(results: dict[str, dict]):
widths = [6] + [10] * len(comps) widths = [6] + [10] * len(comps)
metric = "er_99_norm" metric = "er_99_norm"
print(f"\n--- Per-Layer Effective Rank (99% energy) ---") print("\n--- Per-Layer Effective Rank (99% energy) ---")
print(format_header(["Layer"] + comps, widths)) print(format_header(["Layer"] + comps, widths))
print("-" * sum(widths)) print("-" * sum(widths))
@@ -173,7 +173,7 @@ def print_layer_grid(results: dict[str, dict]):
def print_weight_stats(results: dict[str, dict]): def print_weight_stats(results: dict[str, dict]):
groups = group_by_component(results) groups = group_by_component(results)
widths = [20, 12, 12, 12, 12] widths = [20, 12, 12, 12, 12]
print(f"\n--- Weight Value Statistics ---") print("\n--- Weight Value Statistics ---")
print(format_header(["Component", "Mean", "Std", "Min", "Max"], widths)) print(format_header(["Component", "Mean", "Std", "Min", "Max"], widths))
print("-" * sum(widths)) print("-" * sum(widths))
@@ -265,7 +265,7 @@ def main():
) )
print(f"{'=' * 70}") print(f"{'=' * 70}")
print(f"Loading weights...") print("Loading weights...")
sd = safetensors.torch.load_file(str(weights_path)) sd = safetensors.torch.load_file(str(weights_path))
print(f" {len(sd)} keys loaded") print(f" {len(sd)} keys loaded")
+10 -16
View File
@@ -20,6 +20,7 @@ from typing import Dict, Iterator, List, Optional, Sequence, Tuple
import numpy as np import numpy as np
import torch import torch
import tqdm import tqdm
from datasets import load_dataset
from astrai.inference import InferenceEngine from astrai.inference import InferenceEngine
from astrai.model import AutoModel from astrai.model import AutoModel
@@ -29,9 +30,7 @@ from astrai.tokenize import AutoTokenizer
# Config # Config
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
HUMANEVAL_URL = ( HUMANEVAL_HF_DATASET = "openai/openai_humaneval"
"https://github.com/openai/human-eval/raw/master/data/HumanEval.jsonl.gz"
)
STOP_SEQUENCES = [ STOP_SEQUENCES = [
"\nclass ", "\nclass ",
@@ -64,21 +63,16 @@ class EvalConfig:
problem_indices: Optional[List[int]] = None problem_indices: Optional[List[int]] = None
def download(url: str, path: str): def download(path: str):
if os.path.exists(path): if os.path.exists(path):
return return
import gzip
import urllib.request
os.makedirs(os.path.dirname(path) or ".", exist_ok=True) os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
print(f"Downloading {url} ...") print(f"Downloading HumanEval from HuggingFace ({HUMANEVAL_HF_DATASET}) ...")
tmp = path + ".tmp" ds = load_dataset(HUMANEVAL_HF_DATASET, split="test")
urllib.request.urlretrieve(url, tmp) with open(path, "w", encoding="utf-8") as f:
with gzip.open(tmp, "rb") as f_in: for item in ds:
with open(path, "wb") as f_out: f.write(json.dumps(item, ensure_ascii=False) + "\n")
f_out.write(f_in.read()) print(f" saved {len(ds)} problems to {path}")
os.remove(tmp)
print(f" saved to {path}")
def load_jsonl(path: str) -> List[dict]: def load_jsonl(path: str) -> List[dict]:
@@ -318,7 +312,7 @@ def run_pipeline(cfg: EvalConfig) -> Dict:
with open(cfg.test_only, encoding="utf-8") as f: with open(cfg.test_only, encoding="utf-8") as f:
generated = json.load(f) generated = json.load(f)
else: else:
download(HUMANEVAL_URL, cfg.data_path) download(cfg.data_path)
problems = load_jsonl(cfg.data_path) problems = load_jsonl(cfg.data_path)
if cfg.problem_indices: if cfg.problem_indices:
+12 -18
View File
@@ -26,28 +26,22 @@ import torch.nn.functional as F
import tqdm import tqdm
from astrai.model import AutoModel from astrai.model import AutoModel
from astrai.preprocessing.packing import plan_bfd
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
def _pack_bins(pairs, max_len): def _pack_bins(pairs, max_len):
"""BFD bin packing: pack (c+r) into bins of max total length.""" """BFD bin packing: pack (c+r) into bins of max total length.
indexed = sorted(enumerate(pairs), key=lambda x: -(len(x[1][0]) + len(x[1][1])))
bins = [] Reuses :func:`plan_bfd` so the BFD heuristic stays single-sourced.
lengths = [] """
for orig_idx, (c, r) in indexed: # Treat each pair as a single sequence of length len(c)+len(r) for
size = len(c) + len(r) # planning purposes; plan_bfd works on pure lengths.
best_bin = -1 fake_sequences = [[0] * (len(c) + len(r)) for c, r in pairs]
for bi, rem in enumerate(lengths): plan = plan_bfd(fake_sequences, max_len)
if rem >= size: return [
if best_bin < 0 or rem < lengths[best_bin]: [(i, pairs[i][0], pairs[i][1]) for i in bin_indices] for bin_indices in plan
best_bin = bi ]
if best_bin >= 0:
bins[best_bin].append((orig_idx, c, r))
lengths[best_bin] -= size
else:
bins.append([(orig_idx, c, r)])
lengths.append(max_len - size)
return bins
def _resolve_sentinel_ids(tokenizer, sentinel_text): def _resolve_sentinel_ids(tokenizer, sentinel_text):
+8 -15
View File
@@ -14,21 +14,17 @@ import argparse
import json import json
import os import os
import re import re
import urllib.request
from typing import Callable, Dict, List, Optional from typing import Callable, Dict, List, Optional
import torch import torch
import tqdm import tqdm
from datasets import load_dataset
from astrai.inference import InferenceEngine from astrai.inference import InferenceEngine
from astrai.model import AutoModel from astrai.model import AutoModel
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
IFEVAL_URL = ( IFEVAL_HF_DATASET = "google/IFEval"
"https://raw.githubusercontent.com/google-research/"
"google-research/master/instruction_following_eval/data/input_data.jsonl"
)
CONSTRAINT_VERIFIERS: Dict[str, Callable[[str, dict], bool]] = {} CONSTRAINT_VERIFIERS: Dict[str, Callable[[str, dict], bool]] = {}
@@ -310,15 +306,12 @@ def download_ifeval(data_path: str):
if os.path.exists(data_path): if os.path.exists(data_path):
return return
os.makedirs(os.path.dirname(data_path) or ".", exist_ok=True) os.makedirs(os.path.dirname(data_path) or ".", exist_ok=True)
print(f"Downloading IFEval from {IFEVAL_URL} ...") print(f"Downloading IFEval from HuggingFace ({IFEVAL_HF_DATASET}) ...")
tmp = data_path + ".tmp" ds = load_dataset(IFEVAL_HF_DATASET, split="train")
urllib.request.urlretrieve(IFEVAL_URL, tmp) with open(data_path, "w", encoding="utf-8") as f:
with open(tmp, "rb") as f_in: for item in ds:
content = f_in.read() f.write(json.dumps(item, ensure_ascii=False) + "\n")
with open(data_path, "wb") as f_out: print(f" saved {len(ds)} items to {data_path}")
f_out.write(content)
os.remove(tmp)
print(f" saved to {data_path}")
def load_problems(data_path: str) -> List[dict]: def load_problems(data_path: str) -> List[dict]:
+86 -44
View File
@@ -4,18 +4,18 @@ import argparse
import csv import csv
import json import json
import os import os
import shutil import random
import tarfile from collections import defaultdict
import requests
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
import tqdm import tqdm
from datasets import load_dataset
from astrai.model import AutoModel from astrai.model import AutoModel
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
MMLU_URL = "https://people.eecs.berkeley.edu/~hendrycks/data.tar" MMLU_HF_DATASET = "cais/mmlu"
MMLU_SUBJECTS = [ MMLU_SUBJECTS = [
"abstract_algebra", "abstract_algebra",
"anatomy", "anatomy",
@@ -77,38 +77,40 @@ MMLU_SUBJECTS = [
] ]
def _download_and_extract(url: str, data_dir: str): def _write_subject_csv(data_dir: str, split: str, subject: str, rows: list[dict]):
tar_path = os.path.join(data_dir, "data.tar") split_dir = os.path.join(data_dir, split)
os.makedirs(data_dir, exist_ok=True) os.makedirs(split_dir, exist_ok=True)
print(f"Downloading MMLU data from {url}...") path = os.path.join(split_dir, f"{subject}_{split}.csv")
resp = requests.get(url, stream=True, timeout=300) with open(path, "w", encoding="utf-8", newline="") as f:
resp.raise_for_status() writer = csv.writer(f)
total = int(resp.headers.get("content-length", 0)) for row in rows:
with tqdm.tqdm(total=total, unit="B", unit_scale=True, desc=" Download") as bar: writer.writerow(row)
with open(tar_path, "wb") as f:
for chunk in resp.iter_content(chunk_size=8192):
f.write(chunk)
bar.update(len(chunk))
print("Extracting...")
with tarfile.open(tar_path, "r") as tf:
tf.extractall(data_dir)
os.remove(tar_path)
def download_mmlu(data_dir: str): def download_mmlu(data_dir: str):
_download_and_extract(MMLU_URL, data_dir) print(f"Downloading MMLU from HuggingFace ({MMLU_HF_DATASET}) ...")
src = os.path.join(data_dir, "data") letters = ("A", "B", "C", "D")
if os.path.exists(src): split_map = {"dev": "dev", "val": "validation", "test": "test"}
for item in os.listdir(src): for local_split, hf_split in split_map.items():
src_item = os.path.join(src, item) ds = load_dataset(MMLU_HF_DATASET, "all", split=hf_split)
dst_item = os.path.join(data_dir, item) grouped: dict[str, list[dict]] = defaultdict(list)
if os.path.exists(dst_item): for item in tqdm.tqdm(ds, desc=f" {local_split}", leave=False):
if os.path.isdir(dst_item): subject = item["subject"]
shutil.rmtree(dst_item) choices = item["choices"]
else: ans_letter = letters[item["answer"]]
os.remove(dst_item) grouped[subject].append(
os.rename(src_item, dst_item) [
os.rmdir(src) item["question"],
f"A){choices[0]}",
f"B){choices[1]}",
f"C){choices[2]}",
f"D){choices[3]}",
ans_letter,
]
)
for subject, rows in grouped.items():
_write_subject_csv(data_dir, local_split, subject, rows)
print(f" {local_split}: {len(ds)} items, {len(grouped)} subjects")
print(f"MMLU data saved to {data_dir}") print(f"MMLU data saved to {data_dir}")
@@ -153,19 +155,22 @@ def build_prompt(question: str, choices: dict, subject: str) -> str:
def apply_chat( def apply_chat(
tokenizer, raw_prompt: str, n_shot: int, dev_data: list[dict] | None tokenizer,
raw_prompt: str,
n_shot: int,
dev_data: list[dict] | None,
subject: str = "",
) -> str: ) -> str:
"""Wrap raw MMLU prompt in the model's chat template format. """Wrap raw MMLU prompt in the model's chat template format.
For few-shot, prepend example Q&A pairs as a second user/assistant exchange. For few-shot, prepend example Q&A pairs as user/assistant exchanges.
Few-shot examples use the same subject preamble as the test question to
keep the format consistent.
""" """
messages = [] messages = []
if n_shot > 0 and dev_data: if n_shot > 0 and dev_data:
for item in dev_data[:n_shot]: for item in dev_data[:n_shot]:
q = f"Question: {item['question']}\n" q = build_prompt(item["question"], item, subject)
for k in ("A", "B", "C", "D"):
q += f"{k}. {item[k]}\n"
q += "Answer:"
messages.append({"role": "user", "content": q}) messages.append({"role": "user", "content": q})
messages.append({"role": "assistant", "content": item["answer"]}) messages.append({"role": "assistant", "content": item["answer"]})
messages.append({"role": "user", "content": raw_prompt}) messages.append({"role": "user", "content": raw_prompt})
@@ -180,7 +185,7 @@ def choice_logprob(
choice_text = choice_letter choice_text = choice_letter
choice_ids = tokenizer.encode(choice_text, add_special_tokens=False) choice_ids = tokenizer.encode(choice_text, add_special_tokens=False)
input_ids = context_ids + choice_ids input_ids = context_ids + choice_ids
max_len = model.config.max_len max_len = model.config.max_position_embeddings
if len(input_ids) > max_len: if len(input_ids) > max_len:
overflow = len(input_ids) - max_len overflow = len(input_ids) - max_len
input_ids = input_ids[overflow:] input_ids = input_ids[overflow:]
@@ -201,6 +206,24 @@ def choice_logprob(
return score return score
def _permute_choices(item: dict, rng: random.Random) -> tuple[dict, str]:
"""Shuffle the option order of a question.
Returns ``(permuted_item, new_answer_letter)``. The question text and
the *content* of each choice are unchanged; only which letter (A/B/C/D)
maps to which content is shuffled. This neutralises the model's
positional bias (e.g. always picking B).
"""
letters = ("A", "B", "C", "D")
perm = list(letters)
rng.shuffle(perm)
permuted = {"question": item["question"]}
for new_letter, orig_letter in zip(letters, perm):
permuted[new_letter] = item[orig_letter]
new_answer = letters[perm.index(item["answer"])]
return permuted, new_answer
def evaluate_subject( def evaluate_subject(
model, model,
tokenizer, tokenizer,
@@ -209,18 +232,24 @@ def evaluate_subject(
dev_data: list[dict] | None, dev_data: list[dict] | None,
device: str, device: str,
n_shot: int, n_shot: int,
seed: int = 0,
) -> tuple[float, int, int]: ) -> tuple[float, int, int]:
rng = random.Random(seed) if seed >= 0 else None
correct = 0 correct = 0
total = 0 total = 0
for item in tqdm.tqdm(test_data, desc=f"{subject:40s}", leave=False): for item in tqdm.tqdm(test_data, desc=f"{subject:40s}", leave=False):
raw_prompt = build_prompt(item["question"], item, subject) if rng is not None:
context = apply_chat(tokenizer, raw_prompt, n_shot, dev_data or []) permuted, answer = _permute_choices(item, rng)
else:
permuted, answer = item, item["answer"]
raw_prompt = build_prompt(permuted["question"], permuted, subject)
context = apply_chat(tokenizer, raw_prompt, n_shot, dev_data or [], subject)
context_ids = tokenizer.encode(context) context_ids = tokenizer.encode(context)
scores = { scores = {
c: choice_logprob(model, tokenizer, context_ids, c, device) c: choice_logprob(model, tokenizer, context_ids, c, device)
for c in ("A", "B", "C", "D") for c in ("A", "B", "C", "D")
} }
if max(scores, key=scores.get) == item["answer"]: if max(scores, key=scores.get) == answer:
correct += 1 correct += 1
total += 1 total += 1
return correct / total, correct, total return correct / total, correct, total
@@ -255,6 +284,12 @@ def main():
default="bfloat16" if torch.cuda.is_available() else "float32", default="bfloat16" if torch.cuda.is_available() else "float32",
help="Torch dtype", help="Torch dtype",
) )
parser.add_argument(
"--seed",
type=int,
default=0,
help="Seed for option permutation (0 to enable, -1 to disable)",
)
args = parser.parse_args() args = parser.parse_args()
if args.download or not os.path.exists(args.data_dir): if args.download or not os.path.exists(args.data_dir):
@@ -286,7 +321,14 @@ def main():
test_data = load_csv(test_path) test_data = load_csv(test_path)
acc, corr, tot = evaluate_subject( acc, corr, tot = evaluate_subject(
model, tokenizer, subject, test_data, dev_data, device, args.n_shot model,
tokenizer,
subject,
test_data,
dev_data,
device,
args.n_shot,
seed=args.seed,
) )
results[subject] = {"accuracy": round(acc, 4), "correct": corr, "total": tot} results[subject] = {"accuracy": round(acc, 4), "correct": corr, "total": tot}
total_correct += corr total_correct += corr
+2 -2
View File
@@ -148,7 +148,7 @@ class LossAccumulator:
self.total += sum(losses) self.total += sum(losses)
self.count += len(losses) self.count += len(losses)
if self.stream: if self.stream:
clamped = [min(max(l, 0.0), self._HIST_MAX) for l in losses] clamped = [min(max(v, 0.0), self._HIST_MAX) for v in losses]
idx = torch.tensor(clamped) / self._HIST_MAX * (self._HIST_BINS - 1) idx = torch.tensor(clamped) / self._HIST_MAX * (self._HIST_BINS - 1)
self.hist += torch.bincount( self.hist += torch.bincount(
idx.long().clamp(0, self._HIST_BINS - 1), idx.long().clamp(0, self._HIST_BINS - 1),
@@ -315,7 +315,7 @@ def print_stats(label: str, stats: Dict):
) )
by_type = stats.get("by_token_type", {}) by_type = stats.get("by_token_type", {})
if by_type: if by_type:
print(f"\n by token type:") print("\n by token type:")
print(f" {'type':<12} {'count':>8} {'mean_loss':>10} {'ppl':>8}") print(f" {'type':<12} {'count':>8} {'mean_loss':>10} {'ppl':>8}")
print(f" {'-' * 12} {'-' * 8} {'-' * 10} {'-' * 8}") print(f" {'-' * 12} {'-' * 8} {'-' * 10} {'-' * 8}")
for ttype, s in by_type.items(): for ttype, s in by_type.items():
+1 -1
View File
@@ -15,7 +15,7 @@ Usage::
import argparse import argparse
import json import json
from collections import Counter from collections import Counter
from typing import Dict, List, Tuple from typing import Dict, List
def _tokenize(text: str) -> List[str]: def _tokenize(text: str) -> List[str]:
+12 -12
View File
@@ -119,15 +119,15 @@ class GenerationBenchmark:
dtype=torch.long, dtype=torch.long,
) )
head_dim = self.config.dim // self.config.n_heads head_dim = self.config.hidden_size // self.config.num_attention_heads
max_seq = prompt_length + gen_length max_seq = prompt_length + gen_length
if self.cache_type == "contiguous": if self.cache_type == "contiguous":
cache = ContiguousCache( cache = ContiguousCache(
self.config.n_layers, self.config.num_hidden_layers,
batch_size, batch_size,
max_seq, max_seq,
self.config.n_kv_heads, self.config.num_key_value_heads,
head_dim, head_dim,
self.device, self.device,
self.dtype, self.dtype,
@@ -136,10 +136,10 @@ class GenerationBenchmark:
page_size = 128 page_size = 128
n_pages = (max_seq + page_size - 1) // page_size * batch_size n_pages = (max_seq + page_size - 1) // page_size * batch_size
cache = PageCache( cache = PageCache(
self.config.n_layers, self.config.num_hidden_layers,
n_pages, n_pages,
page_size, page_size,
self.config.n_kv_heads, self.config.num_key_value_heads,
head_dim, head_dim,
self.device, self.device,
self.dtype, self.dtype,
@@ -262,13 +262,13 @@ if __name__ == "__main__":
config = AutoRegressiveLMConfig( config = AutoRegressiveLMConfig(
vocab_size=10000, vocab_size=10000,
dim=1536, hidden_size=1536,
n_heads=24, num_attention_heads=24,
n_kv_heads=4, num_key_value_heads=4,
dim_ffn=6912, intermediate_size=6912,
max_len=2048, max_position_embeddings=2048,
n_layers=24, num_hidden_layers=24,
norm_eps=1e-5, rms_norm_eps=1e-5,
) )
benchmark = GenerationBenchmark( benchmark = GenerationBenchmark(
+94 -18
View File
@@ -1,8 +1,10 @@
import argparse import argparse
import json import json
import time
from typing import Optional from typing import Optional
import torch import torch
from tqdm import tqdm
from astrai.inference import InferenceEngine from astrai.inference import InferenceEngine
from astrai.model import AutoModel from astrai.model import AutoModel
@@ -20,55 +22,102 @@ def processor(
response_key: str, response_key: str,
max_tokens: Optional[int], max_tokens: Optional[int],
batch_size: int, batch_size: int,
num_samples: int = 1,
cache_len: int = 2048,
frequency_penalty: float = 0.0,
rep_window: int = 64,
): ):
# Load model and tokenizer print(f"Loading model from {param_path} ...")
t0 = time.time()
model = AutoModel.from_pretrained(param_path) model = AutoModel.from_pretrained(param_path)
tokenizer = AutoTokenizer.from_pretrained(param_path) tokenizer = AutoTokenizer.from_pretrained(param_path)
model.to(device="cuda", dtype=torch.bfloat16) model.to(device="cuda", dtype=torch.bfloat16)
print(f" model loaded in {time.time() - t0:.1f}s")
# Create inference engine
engine = InferenceEngine( engine = InferenceEngine(
model=model, tokenizer=tokenizer, max_batch_size=batch_size model=model,
tokenizer=tokenizer,
max_batch_size=batch_size * num_samples,
max_seq_len=cache_len,
max_prompt_len=cache_len,
) )
print(f"Reading {input_json_file} ...")
with open(input_json_file, "r", encoding="utf-8") as f: with open(input_json_file, "r", encoding="utf-8") as f:
input_data = [json.loads(line) for line in f] input_data = [json.loads(line) for line in f]
# Check input format: chat messages or raw text
if input_data and "messages" in input_data[0]: if input_data and "messages" in input_data[0]:
# Chat format: [{"messages": [...]}]
prompts = [ prompts = [
tokenizer.apply_chat_template(item["messages"], tokenize=False) tokenizer.apply_chat_template(item["messages"], tokenize=False)
for item in input_data for item in input_data
] ]
else: else:
# Raw text format: [{"question": "..."}]
prompts = [item[question_key] for item in input_data] prompts = [item[question_key] for item in input_data]
print(f" {len(prompts)} prompts loaded\n")
# Use provided max_tokens or default to model config max_len
if max_tokens is None: if max_tokens is None:
max_tokens = model.config.max_len max_tokens = model.config.max_position_embeddings
# Generate responses (batch) chunk_size = max(1, batch_size)
responses = engine.generate(
prompt=prompts, with open(output_json_file, "w", encoding="utf-8") as f:
pbar = tqdm(
total=len(prompts) * num_samples,
unit="gen",
desc=f" Generating ({num_samples}x/prompt)",
)
for chunk_start in range(0, len(prompts), chunk_size):
chunk = prompts[chunk_start : chunk_start + chunk_size]
if num_samples > 1:
chunk_expanded = [p for p in chunk for _ in range(num_samples)]
resp_chunk = engine.generate(
prompt=chunk_expanded,
stream=False, stream=False,
max_tokens=max_tokens, max_tokens=max_tokens,
temperature=temperature, temperature=temperature,
top_p=top_p, top_p=top_p,
top_k=top_k, top_k=top_k,
frequency_penalty=frequency_penalty,
rep_window=rep_window,
)
resp_chunk = [
resp_chunk[i * num_samples : (i + 1) * num_samples]
for i in range(len(chunk))
]
else:
resp_chunk = engine.generate(
prompt=chunk,
stream=False,
max_tokens=max_tokens,
temperature=temperature,
top_p=top_p,
top_k=top_k,
frequency_penalty=frequency_penalty,
rep_window=rep_window,
) )
# Write results for i, prompt in enumerate(chunk):
with open(output_json_file, "w", encoding="utf-8") as f:
for prompt, response in zip(prompts, responses):
if input_data and "messages" in input_data[0]: if input_data and "messages" in input_data[0]:
output_item = {"response": response} orig = input_data[chunk_start + i]
output_item = {**orig, response_key: resp_chunk[i]}
else: else:
output_item = {question_key: prompt, response_key: response} output_item = {
question_key: prompt,
response_key: resp_chunk[i],
}
f.write(json.dumps(output_item, ensure_ascii=False) + "\n") f.write(json.dumps(output_item, ensure_ascii=False) + "\n")
# Cleanup pbar.update(len(chunk) * num_samples)
pbar.close()
elapsed = time.time() - t0
print(
f"\nDone! {len(prompts)} prompts x {num_samples} samples -> {output_json_file}"
)
print(f"Total time: {elapsed:.1f}s ({elapsed / len(prompts):.2f}s/prompt)")
engine.shutdown() engine.shutdown()
@@ -126,11 +175,38 @@ if __name__ == "__main__":
default=1, default=1,
help="Batch size for generating responses (default: 1).", help="Batch size for generating responses (default: 1).",
) )
parser.add_argument(
"--num_samples",
type=int,
default=1,
help="Number of responses per prompt (expands batch internally, default: 1).",
)
parser.add_argument( parser.add_argument(
"--max_tokens", "--max_tokens",
type=int, type=int,
default=None, default=None,
help="Maximum tokens to generate (default: model config max_len).", help=(
"Maximum tokens to generate "
"(default: model config max_position_embeddings)."
),
)
parser.add_argument(
"--cache_len",
type=int,
default=2048,
help="KV cache & prompt truncation length (default: 2048, lower = less memory).",
)
parser.add_argument(
"--frequency_penalty",
type=float,
default=0.0,
help="Frequency penalty to reduce repetition (default: 0.0, try 0.5-1.0).",
)
parser.add_argument(
"--rep_window",
type=int,
default=64,
help="Window size for frequency penalty (default: 64).",
) )
args = parser.parse_args() args = parser.parse_args()
+10
View File
@@ -22,9 +22,19 @@ def main():
default="params", default="params",
help="Path to tokenizer directory (default: params)", help="Path to tokenizer directory (default: params)",
) )
parser.add_argument(
"--batch_size",
type=int,
default=None,
help="Number of records tokenized together (default: config value)",
)
args = parser.parse_args() args = parser.parse_args()
config = PipelineConfig.from_file(args.config) config = PipelineConfig.from_file(args.config)
if args.batch_size is not None:
if args.batch_size < 1:
parser.error("--batch_size must be at least 1")
config.preprocessing.batch_size = args.batch_size
Pipeline( Pipeline(
config=config, config=config,
+81 -13
View File
@@ -1,17 +1,18 @@
import argparse import argparse
import os import os
from functools import partial from functools import partial
from typing import Any, Dict from typing import Any, Callable, Dict, Optional
import torch import torch
import torch.optim as optim import torch.optim as optim
from torch import Tensor, nn from torch import Tensor, nn
from astrai.config import AutoRegressiveLMConfig, TrainConfig from astrai.config import AutoRegressiveLMConfig, TrainConfig
from astrai.dataset import DatasetFactory from astrai.dataset import DatasetFactory, dpo_collate_fn, grpo_collate_fn
from astrai.model import AutoRegressiveLM from astrai.model import AutoRegressiveLM
from astrai.model.components.decoder_block import DecoderBlock from astrai.model.components.decoder_block import DecoderBlock
from astrai.trainer import SchedulerFactory, Trainer from astrai.trainer import SchedulerFactory, Trainer
from astrai.trainer.rollout import BaseRewardModel
class MuonMix(optim.Optimizer): class MuonMix(optim.Optimizer):
@@ -101,7 +102,7 @@ def parse_args() -> argparse.Namespace:
"--train_type", "--train_type",
type=str, type=str,
required=True, required=True,
choices=["seq", "sft", "dpo", "grpo"], choices=["seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"],
help="Train type.", help="Train type.",
) )
parser.add_argument( parser.add_argument(
@@ -149,7 +150,7 @@ def parse_args() -> argparse.Namespace:
"--max_grad_norm", "--max_grad_norm",
type=float, type=float,
default=1.0, default=1.0,
help="Max gradient norm for clipping.", help="Max gradient norm for clipping. None disables clipping.",
) )
parser.add_argument( parser.add_argument(
"--weight_decay", "--weight_decay",
@@ -217,6 +218,39 @@ def parse_args() -> argparse.Namespace:
default=0.0, default=0.0,
help="cross_entropy function label smoothing parameter", help="cross_entropy function label smoothing parameter",
) )
# online rollout
parser.add_argument(
"--rollout_interval",
type=int,
default=512,
help="Number of optimizer steps between online rollouts.",
)
parser.add_argument(
"--rollout_temperature",
type=float,
default=0.7,
help="Sampling temperature for online rollout.",
)
parser.add_argument(
"--rollout_top_k",
type=int,
default=0,
help="Top-k filtering for online rollout (0=disable).",
)
parser.add_argument(
"--rollout_top_p",
type=float,
default=0.9,
help="Top-p (nucleus) filtering for online rollout.",
)
parser.add_argument(
"--rollout_max_tokens",
type=int,
default=1024,
help="Maximum generated tokens per response in rollout.",
)
parser.add_argument( parser.add_argument(
"--gradient_checkpointing", "--gradient_checkpointing",
action=argparse.BooleanOptionalAction, action=argparse.BooleanOptionalAction,
@@ -293,8 +327,8 @@ def parse_args() -> argparse.Namespace:
"--parallel_mode", "--parallel_mode",
type=str, type=str,
default="none", default="none",
choices=["none", "ddp", "fsdp"], choices=["none", "ddp", "fsdp", "fsdp2"],
help="Parallel training strategy (none, ddp, fsdp).", help="Parallel training strategy (none, ddp, fsdp, fsdp2).",
) )
parser.add_argument( parser.add_argument(
"--device_type", type=str, default="cuda", help="Device type to use." "--device_type", type=str, default="cuda", help="Device type to use."
@@ -428,10 +462,19 @@ def train(
decay_steps: int, decay_steps: int,
**kwargs, **kwargs,
): ):
assert train_type in ["seq", "sft", "dpo", "grpo"] assert train_type in [
"seq",
"sft",
"dpo",
"grpo",
"online_grpo",
"online_dpo",
]
assert os.path.exists(param_path) assert os.path.exists(param_path)
if nprocs > 1 and parallel_mode == "none": if nprocs > 1 and parallel_mode == "none":
raise ValueError("--nprocs > 1 requires --parallel_mode to be 'ddp' or 'fsdp'") raise ValueError(
"--nprocs > 1 requires --parallel_mode to be 'ddp', 'fsdp', or 'fsdp2'"
)
# Load config # Load config
config_path = os.path.join(param_path, "config.json") config_path = os.path.join(param_path, "config.json")
@@ -439,7 +482,7 @@ def train(
config.neftune_alpha = neftune_alpha config.neftune_alpha = neftune_alpha
if window_size is None: if window_size is None:
window_size = config.max_len window_size = config.max_position_embeddings
strategy_kwargs = { strategy_kwargs = {
"beta": kwargs.pop("dpo_beta"), "beta": kwargs.pop("dpo_beta"),
@@ -449,10 +492,19 @@ def train(
"group_size": kwargs.pop("group_size"), "group_size": kwargs.pop("group_size"),
} }
executor_kwargs = { rollout_interval = kwargs.pop("rollout_interval", 512)
"gradient_as_bucket_view": True, rollout_temperature = kwargs.pop("rollout_temperature", 0.7)
"broadcast_buffers": False, rollout_top_k = kwargs.pop("rollout_top_k", 0)
} rollout_top_p = kwargs.pop("rollout_top_p", 0.9)
rollout_max_tokens = kwargs.pop("rollout_max_tokens", 1024)
reward_model_fn: Optional[Callable[[], BaseRewardModel]] = None
executor_kwargs = {}
if parallel_mode == "ddp":
executor_kwargs.update(
gradient_as_bucket_view=True,
broadcast_buffers=False,
)
model_fn = partial(create_model, config) model_fn = partial(create_model, config)
dataset = DatasetFactory.load( dataset = DatasetFactory.load(
@@ -460,6 +512,7 @@ def train(
load_path=data_root_path, load_path=data_root_path,
window_size=window_size, window_size=window_size,
stride=stride, stride=stride,
tokenizer_path=param_path,
) )
optimizer_fn = partial( optimizer_fn = partial(
@@ -504,6 +557,14 @@ def train(
grad_ckpt_modules = [DecoderBlock] if gradient_checkpointing else [] grad_ckpt_modules = [DecoderBlock] if gradient_checkpointing else []
collate_fn = None
if train_type == "dpo":
collate_fn = dpo_collate_fn
elif train_type == "grpo":
collate_fn = grpo_collate_fn
elif train_type in ("online_grpo", "online_dpo"):
collate_fn = None
train_config = TrainConfig( train_config = TrainConfig(
model_fn=model_fn, model_fn=model_fn,
strategy=train_type, strategy=train_type,
@@ -536,6 +597,13 @@ def train(
executor_kwargs=executor_kwargs, executor_kwargs=executor_kwargs,
extra_kwargs=strategy_kwargs, extra_kwargs=strategy_kwargs,
neftune_alpha=neftune_alpha, neftune_alpha=neftune_alpha,
collate_fn=collate_fn,
rollout_interval=rollout_interval,
rollout_temperature=rollout_temperature,
rollout_top_k=rollout_top_k,
rollout_top_p=rollout_top_p,
rollout_max_tokens=rollout_max_tokens,
reward_model_fn=reward_model_fn,
) )
trainer = Trainer(train_config) trainer = Trainer(train_config)
+14 -14
View File
@@ -107,13 +107,13 @@ def test_model():
"""Session-scoped small AutoRegressiveLM model, created once.""" """Session-scoped small AutoRegressiveLM model, created once."""
config = AutoRegressiveLMConfig( config = AutoRegressiveLMConfig(
vocab_size=1000, vocab_size=1000,
dim=8, hidden_size=8,
n_heads=2, num_attention_heads=2,
n_kv_heads=1, num_key_value_heads=1,
dim_ffn=16, intermediate_size=16,
max_len=64, max_position_embeddings=64,
n_layers=2, num_hidden_layers=2,
norm_eps=1e-5, rms_norm_eps=1e-5,
) )
device = "cuda" if torch.cuda.is_available() else "cpu" device = "cuda" if torch.cuda.is_available() else "cpu"
model = AutoRegressiveLM(config).to(device=device) model = AutoRegressiveLM(config).to(device=device)
@@ -137,13 +137,13 @@ def base_test_env(test_model, test_tokenizer):
json.dump( json.dump(
{ {
"vocab_size": 1000, "vocab_size": 1000,
"dim": 8, "hidden_size": 8,
"n_heads": 2, "num_attention_heads": 2,
"n_kv_heads": 1, "num_key_value_heads": 1,
"dim_ffn": 16, "intermediate_size": 16,
"max_len": 64, "max_position_embeddings": 64,
"n_layers": 2, "num_hidden_layers": 2,
"norm_eps": 1e-5, "rms_norm_eps": 1e-5,
}, },
f, f,
) )
+350 -51
View File
@@ -1,14 +1,16 @@
import json import json
import os import os
import tempfile
import numpy as np import numpy as np
import pytest import pytest
import torch import torch
from astrai.config.preprocess_config import PipelineConfig from astrai.config.preprocess_config import PipelineConfig
from astrai.dataset.dataset import DatasetFactory, SEQDataset from astrai.dataset.dataset import DatasetFactory, dpo_tokenize
from astrai.dataset.storage import ( from astrai.dataset.storage import (
H5Store, H5Store,
JsonlStore,
StoreFactory, StoreFactory,
detect_format, detect_format,
) )
@@ -56,6 +58,13 @@ def _write_jsonl_dataset(test_dir, tokenizer_path, records, config_overrides=Non
return data_dir return data_dir
def _fake_fetch_record(self, idx, keys):
"""FakeStore.fetch_record matching real Store semantics."""
if isinstance(keys, str):
return self._data[keys][idx]
return {k: self._data[k][idx] for k in keys}
def _make_seq_dataset( def _make_seq_dataset(
test_dir, name="data", seq_length=200, train_type="seq", data=None, **load_kwargs test_dir, name="data", seq_length=200, train_type="seq", data=None, **load_kwargs
): ):
@@ -109,7 +118,7 @@ def test_dpo_strategy_with_random_data(base_test_env):
) )
assert dpo_dataset is not None assert dpo_dataset is not None
assert dpo_dataset.storage is not None assert dpo_dataset.store is not None
assert len(dpo_dataset) > 0 assert len(dpo_dataset) > 0
# Test that we can get DPO items without errors # Test that we can get DPO items without errors
@@ -138,7 +147,7 @@ def test_sft_dataset_with_random_data(base_test_env):
) )
assert sft_dataset is not None assert sft_dataset is not None
assert sft_dataset.storage is not None assert sft_dataset.store is not None
assert len(sft_dataset) > 0 assert len(sft_dataset) > 0
# Test that we can get SFT items without errors # Test that we can get SFT items without errors
@@ -169,39 +178,37 @@ def test_dataset_with_custom_stride(base_test_env):
assert len(dataset) > len(default_stride_dataset) assert len(dataset) > len(default_stride_dataset)
def test_dataset_count_property(base_test_env): def test_dataset_token_count_property(base_test_env):
"""dataset.token_count exposes the raw stream token length."""
test_dir = base_test_env["test_dir"] test_dir = base_test_env["test_dir"]
dataset = _make_seq_dataset(test_dir, "count_test_data") dataset = _make_seq_dataset(test_dir, "count_test_data")
assert dataset.count == 200 assert dataset.token_count == 200
assert dataset.count > len(dataset) assert dataset.token_count > len(dataset)
assert len(dataset) == (200 - 1 - 64) // 64 + 1 assert len(dataset) == (200 - 1 - 64) // 64 + 1
def test_empty_dataset_count():
"""Test count returns 0 when no data is loaded"""
dataset = SEQDataset(window_size=64, stride=32)
assert dataset.count == 0
assert dataset.keys == []
def test_dataset_too_short_for_window(base_test_env): def test_dataset_too_short_for_window(base_test_env):
test_dir = base_test_env["test_dir"] test_dir = base_test_env["test_dir"]
dataset = _make_seq_dataset(test_dir, "short", seq_length=30) dataset = _make_seq_dataset(test_dir, "short", seq_length=30)
assert len(dataset) == 0 assert len(dataset) == 0
assert dataset.count == 30 assert dataset.token_count == 30
def test_unloaded_dataset_getitem_raises(): def test_unloaded_sample_window_raises():
"""__getitem__ without load() should fail clearly""" """Store.sample_window before load raises RuntimeError."""
dataset = SEQDataset(window_size=64, stride=32) from astrai.dataset.storage import H5Store
with pytest.raises(RuntimeError, match="not loaded"):
dataset.get_index(0) store = H5Store(window_size=64, stride=64)
with pytest.raises(IndexError, match="Data too short"):
store.sample_window(0)
def test_unloaded_dataset_len(): def test_unloaded_dataset_len():
"""__len__ without load() returns 0""" """__len__ on a store with no data returns 0."""
dataset = SEQDataset(window_size=64, stride=32) from astrai.dataset.storage import H5Store
assert len(dataset) == 0
store = H5Store(window_size=64, stride=64)
assert len(store) == 0
def test_store_unloaded_len(): def test_store_unloaded_len():
@@ -214,7 +221,7 @@ def test_store_unloaded_len():
def test_store_fetch_begin_equals_end(base_test_env): def test_store_fetch_begin_equals_end(base_test_env):
test_dir = base_test_env["test_dir"] test_dir = base_test_env["test_dir"]
dataset = _make_seq_dataset(test_dir, "empty_fetch", seq_length=100, window_size=32) dataset = _make_seq_dataset(test_dir, "empty_fetch", seq_length=100, window_size=32)
result = dataset.storage.fetch(10, 10, "sequence") result = dataset.store.fetch(10, 10, "sequence")
assert result.numel() == 0 assert result.numel() == 0
@@ -264,7 +271,7 @@ def test_store_multi_segment_concat(base_test_env):
store = StoreFactory.create("h5") store = StoreFactory.create("h5")
store.load(data_dir) store.load(data_dir)
assert len(store) == 9 assert store.token_count == 9
result = store.fetch(2, 7, "sequence") result = store.fetch(2, 7, "sequence")
assert result.tolist() == [3, 4, 5, 6, 7] assert result.tolist() == [3, 4, 5, 6, 7]
@@ -293,7 +300,9 @@ def test_mmap_store_load_and_fetch(base_test_env):
store = StoreFactory.create("bin") store = StoreFactory.create("bin")
store.load(test_dir) store.load(test_dir)
assert len(store) == 200 assert store.token_count == 200
assert store.num_records == 0
assert len(store) == 0 # no window configured, no records → 0 samples
assert "sequence" in store.keys assert "sequence" in store.keys
result = store.fetch(10, 20, "sequence") result = store.fetch(10, 20, "sequence")
@@ -306,23 +315,26 @@ def test_mmap_dataset_load(base_test_env):
save_bin(test_dir, data) save_bin(test_dir, data)
dataset = DatasetFactory.load("seq", test_dir, window_size=64) dataset = DatasetFactory.load("seq", test_dir, window_size=64)
assert len(dataset) > 0 assert len(dataset) > 0
assert dataset.count == 200 assert dataset.token_count == 200
assert dataset[0]["input_ids"].shape[0] == 64 assert dataset[0]["input_ids"].shape[0] == 64
def test_normalize_empty_key(): def test_normalize_empty_key():
"""_normalize with empty tensor list does not crash""" """_normalize with empty tensor list does not crash."""
store = H5Store() store = H5Store()
store._normalize({"sequence": []}) store._normalize({"sequence": []})
assert len(store) == 0 assert len(store) == 0
assert store.num_records == 0 # empty key forces num_records=0
assert store.keys == ["sequence"] assert store.keys == ["sequence"]
def test_normalize_mixed_empty_key(): def test_normalize_mixed_empty_key():
"""_normalize with empty + non-empty keys returns min=0""" """_normalize with empty + non-empty keys returns min=0 records."""
store = H5Store() store = H5Store()
store._normalize({"sequence": [torch.tensor([1, 2, 3])], "loss_mask": []}) store._normalize({"sequence": [torch.tensor([1, 2, 3])], "loss_mask": []})
assert len(store) == 0 assert len(store) == 0
assert store.num_records == 0
assert store.token_count == 0 # min() over keys
assert set(store.keys) == {"sequence", "loss_mask"} assert set(store.keys) == {"sequence", "loss_mask"}
@@ -330,14 +342,14 @@ def test_grpo_dataset_dtype(base_test_env):
"""GRPO dataset returns correct dtypes for per-record structured data.""" """GRPO dataset returns correct dtypes for per-record structured data."""
from astrai.dataset.dataset import GRPODataset from astrai.dataset.dataset import GRPODataset
test_dir = base_test_env["test_dir"]
G = 4 G = 4
dataset = GRPODataset() store = type(
dataset.storage = type(
"FakeStore", "FakeStore",
(), (),
{ {
"keys": ["prompts", "responses", "masks", "rewards"], "keys": ["prompts", "responses", "masks", "rewards"],
"num_records": 1,
"token_count": 0,
"_data": { "_data": {
"prompts": [torch.randint(0, 100, (10,), dtype=torch.int32)], "prompts": [torch.randint(0, 100, (10,), dtype=torch.int32)],
"responses": [ "responses": [
@@ -346,9 +358,11 @@ def test_grpo_dataset_dtype(base_test_env):
"masks": [[torch.ones(5, dtype=torch.int32) for _ in range(G)]], "masks": [[torch.ones(5, dtype=torch.int32) for _ in range(G)]],
"rewards": [torch.rand(G, dtype=torch.float32)], "rewards": [torch.rand(G, dtype=torch.float32)],
}, },
"fetch_record": _fake_fetch_record,
"__len__": lambda self: self.num_records,
}, },
)() )()
dataset._build_records() dataset = GRPODataset(store=store)
item = dataset[0] item = dataset[0]
assert item["prompts"].dtype == torch.long assert item["prompts"].dtype == torch.long
@@ -361,25 +375,27 @@ def test_grpo_dataset_load(base_test_env):
"""GRPO dataset loads record-structured data with per-response boundaries.""" """GRPO dataset loads record-structured data with per-response boundaries."""
from astrai.dataset.dataset import GRPODataset from astrai.dataset.dataset import GRPODataset
test_dir = base_test_env["test_dir"]
G = 3 G = 3
prompt_len = 8 prompt_len = 8
resp_lens = [5, 7, 4] resp_lens = [5, 7, 4]
dataset = GRPODataset() store = type(
dataset.storage = type(
"FakeStore", "FakeStore",
(), (),
{ {
"keys": ["prompts", "responses", "masks", "rewards"], "keys": ["prompts", "responses", "masks", "rewards"],
"num_records": 1,
"token_count": 0,
"_data": { "_data": {
"prompts": [torch.randint(0, 100, (prompt_len,))], "prompts": [torch.randint(0, 100, (prompt_len,))],
"responses": [[torch.randint(0, 100, (rl,)) for rl in resp_lens]], "responses": [[torch.randint(0, 100, (rl,)) for rl in resp_lens]],
"masks": [[torch.ones(rl, dtype=torch.int64) for rl in resp_lens]], "masks": [[torch.ones(rl, dtype=torch.int64) for rl in resp_lens]],
"rewards": [torch.tensor([0.9, 0.3, 0.7], dtype=torch.float32)], "rewards": [torch.tensor([0.9, 0.3, 0.7], dtype=torch.float32)],
}, },
"fetch_record": _fake_fetch_record,
"__len__": lambda self: self.num_records,
}, },
)() )()
dataset._build_records() dataset = GRPODataset(store=store)
assert len(dataset) == 1 assert len(dataset) == 1
item = dataset[0] item = dataset[0]
@@ -447,16 +463,17 @@ def test_dataset_load_explicit_storage_type(base_test_env):
test_dir = base_test_env["test_dir"] test_dir = base_test_env["test_dir"]
dataset = _make_seq_dataset(test_dir, "explicit", storage_type="h5") dataset = _make_seq_dataset(test_dir, "explicit", storage_type="h5")
assert len(dataset) > 0 assert len(dataset) > 0
assert dataset.count == 200 assert dataset.token_count == 200
def _write_json_dataset(test_dir, tokenizer_path, records, config_overrides=None): def _write_json_dataset(test_dir, tokenizer_path, records, config_overrides=None):
"""Write JSON (not JSONL) dataset — array of objects.""" """Write JSONL dataset — one JSON object per line."""
data_dir = os.path.join(test_dir, "json_data") data_dir = os.path.join(test_dir, "json_data")
os.makedirs(data_dir, exist_ok=True) os.makedirs(data_dir, exist_ok=True)
with open(os.path.join(data_dir, "data.json"), "w", encoding="utf-8") as f: with open(os.path.join(data_dir, "data.jsonl"), "w", encoding="utf-8") as f:
json.dump(records, f, ensure_ascii=False) for rec in records:
f.write(json.dumps(rec, ensure_ascii=False) + "\n")
config = { config = {
"tokenizer_path": tokenizer_path, "tokenizer_path": tokenizer_path,
@@ -535,7 +552,7 @@ def test_json_store_no_tokenizer_path(base_test_env):
# Save tokenizer files directly in the dataset directory # Save tokenizer files directly in the dataset directory
tokenizer.save_pretrained(data_dir) tokenizer.save_pretrained(data_dir)
# Write .json data # Write .jsonl data
records = [ records = [
{ {
"messages": [ "messages": [
@@ -544,8 +561,9 @@ def test_json_store_no_tokenizer_path(base_test_env):
] ]
} }
] ]
with open(os.path.join(data_dir, "data.json"), "w", encoding="utf-8") as f: with open(os.path.join(data_dir, "data.jsonl"), "w", encoding="utf-8") as f:
json.dump(records, f, ensure_ascii=False) for rec in records:
f.write(json.dumps(rec, ensure_ascii=False) + "\n")
# dataset_config.json WITHOUT tokenizer_path # dataset_config.json WITHOUT tokenizer_path
config = { config = {
@@ -636,6 +654,92 @@ def test_jsonl_store_sft(base_test_env):
assert item["loss_mask"].dtype == torch.bool assert item["loss_mask"].dtype == torch.bool
def test_sft_jsonl_default_messages_config(base_test_env):
"""SFT loads a chat-style JSONL dir with no dataset_config.json.
Falls back to the built-in messages config: every role except
``assistant`` is masked, loss on assistant only.
"""
test_dir = base_test_env["test_dir"]
tokenizer = base_test_env["tokenizer"]
tokenizer.set_chat_template(
"{% for message in messages %}{{ message['role'] }}:{{ message['content'] }}\n{% endfor %}"
)
tokenizer_path = _save_test_tokenizer(test_dir, tokenizer)
data_dir = os.path.join(test_dir, "jsonl_data")
os.makedirs(data_dir, exist_ok=True)
records = [
{
"messages": [
{"role": "system", "content": "sys"},
{"role": "user", "content": "hi"},
{"role": "assistant", "content": "hello"},
]
},
{
"messages": [
{"role": "user", "content": "bye"},
{"role": "assistant", "content": "see you"},
]
},
]
with open(os.path.join(data_dir, "data.jsonl"), "w", encoding="utf-8") as f:
for record in records:
f.write(json.dumps(record, ensure_ascii=False) + "\n")
dataset = DatasetFactory.load(
"sft", data_dir, window_size=8, tokenizer_path=tokenizer_path
)
assert "sequence" in dataset.keys
assert "loss_mask" in dataset.keys
assert "position_ids" in dataset.keys
assert len(dataset) > 0
item = dataset[0]
assert "input_ids" in item
assert "target_ids" in item
assert "loss_mask" in item
assert "position_ids" in item
assert item["loss_mask"].dtype == torch.bool
def test_sft_jsonl_explicit_config_takes_priority(base_test_env):
"""When dataset_config.json exists, it overrides the default messages config."""
test_dir = base_test_env["test_dir"]
tokenizer = base_test_env["tokenizer"]
tokenizer.set_chat_template(
"{% for message in messages %}{{ message['role'] }}:{{ message['content'] }}\n{% endfor %}"
)
tokenizer_path = _save_test_tokenizer(test_dir, tokenizer)
data_dir = _write_jsonl_dataset(
test_dir,
tokenizer_path,
[
{
"messages": [
{"role": "user", "content": "hi"},
{"role": "assistant", "content": "hello"},
]
}
],
config_overrides={
"input": {
"sections": [{"field": "messages", "action": "$role", "template": True}]
},
"mask": {"user": "mask", "assistant": "train"},
"mask_default": "mask",
"preprocessing": {"max_seq_len": 128},
"output": {"position_ids_mode": "doc_reset"},
},
)
dataset = DatasetFactory.load(
"sft", data_dir, window_size=8, tokenizer_path=tokenizer_path
)
assert "sequence" in dataset.keys
assert "loss_mask" in dataset.keys
def test_jsonl_store_pipeline_config_roundtrip(base_test_env): def test_jsonl_store_pipeline_config_roundtrip(base_test_env):
test_dir = base_test_env["test_dir"] test_dir = base_test_env["test_dir"]
config_path = os.path.join(test_dir, "dataset_config.json") config_path = os.path.join(test_dir, "dataset_config.json")
@@ -720,7 +824,7 @@ def test_grpo_builder_preserves_response_boundaries(base_test_env):
from tests.data.conftest import make_grpo_no_template_config from tests.data.conftest import make_grpo_no_template_config
tokenizer = base_test_env["tokenizer"] tokenizer = base_test_env["tokenizer"]
tokenizer_path = _save_test_tokenizer(base_test_env["test_dir"], tokenizer) _save_test_tokenizer(base_test_env["test_dir"], tokenizer)
builder = SectionedMaskBuilder() builder = SectionedMaskBuilder()
config = make_grpo_no_template_config() config = make_grpo_no_template_config()
@@ -833,18 +937,19 @@ def test_grpo_collate_variable_lengths():
assert result["masks"].shape == (2, 2, 4) assert result["masks"].shape == (2, 2, 4)
assert result["rewards"].shape == (2, 2) assert result["rewards"].shape == (2, 2)
# Check padding: item 1 prompt is length 2, padded to 3 # Prompts are left-padded so each response follows its real prompt tokens.
assert result["prompts"][1, 2] == 0 assert torch.equal(result["prompts"][1], torch.tensor([0, 10, 11]))
assert torch.equal(result["prompt_mask"][1], torch.tensor([False, True, True]))
# Check response content: item 0, response 0 is [4,5] padded to 4 # Check response content: item 0, response 0 is [4,5] padded to 4
assert result["responses"][0, 0, 0] == 4 assert result["responses"][0, 0, 0] == 4
assert result["responses"][0, 0, 1] == 5 assert result["responses"][0, 0, 1] == 5
assert result["responses"][0, 0, 2] == 0 # padded assert result["responses"][0, 0, 2] == 0 # padded
assert result["masks"][0, 0, 2] == False # padded assert not result["masks"][0, 0, 2] # padded
# Check response content: item 0, response 1 is [6,7,8,9] no padding # Check response content: item 0, response 1 is [6,7,8,9] no padding
assert result["responses"][0, 1, 3] == 9 assert result["responses"][0, 1, 3] == 9
assert result["masks"][0, 1, 3] == True assert result["masks"][0, 1, 3]
def test_grpo_multiple_records(base_test_env): def test_grpo_multiple_records(base_test_env):
@@ -858,12 +963,13 @@ def test_grpo_multiple_records(base_test_env):
[torch.randint(0, 100, (np.random.randint(3, 8),)) for _ in range(G)] [torch.randint(0, 100, (np.random.randint(3, 8),)) for _ in range(G)]
for _ in range(n_records) for _ in range(n_records)
] ]
dataset = GRPODataset() store = type(
dataset.storage = type(
"FakeStore", "FakeStore",
(), (),
{ {
"keys": ["prompts", "responses", "masks", "rewards"], "keys": ["prompts", "responses", "masks", "rewards"],
"num_records": n_records,
"token_count": 0,
"_data": { "_data": {
"prompts": [torch.randint(0, 100, (10,)) for _ in range(n_records)], "prompts": [torch.randint(0, 100, (10,)) for _ in range(n_records)],
"responses": dummy_responses, "responses": dummy_responses,
@@ -875,9 +981,11 @@ def test_grpo_multiple_records(base_test_env):
torch.rand(G, dtype=torch.float32) for _ in range(n_records) torch.rand(G, dtype=torch.float32) for _ in range(n_records)
], ],
}, },
"fetch_record": _fake_fetch_record,
"__len__": lambda self: self.num_records,
}, },
)() )()
dataset._build_records() dataset = GRPODataset(store=store)
assert len(dataset) == n_records assert len(dataset) == n_records
@@ -888,3 +996,194 @@ def test_grpo_multiple_records(base_test_env):
assert item["rewards"].shape == (G,) assert item["rewards"].shape == (G,)
for g in range(G): for g in range(G):
assert item["responses"][g].shape == item["masks"][g].shape assert item["responses"][g].shape == item["masks"][g].shape
def _write_dpo_jsonl(test_dir, records):
"""Write a raw DPO JSONL file (no dataset_config.json)."""
path = os.path.join(test_dir, "dpo.jsonl")
with open(path, "w", encoding="utf-8") as f:
for rec in records:
f.write(json.dumps(rec, ensure_ascii=False) + "\n")
return path
def test_dpo_tokenize_pure_function():
"""dpo_tokenize returns flat lists with correct mask alignment."""
class FakeTokenizer:
def apply_chat_template(
self, messages, tokenize=True, add_generation_prompt=True
):
ids = []
for m in messages:
ids.append(len(m["content"]))
ids.append(-1)
if add_generation_prompt:
ids.append(99)
return ids
record = {"prompt": "ab", "chosen": "xyz", "rejected": "w"}
result = dpo_tokenize(record, FakeTokenizer(), max_len=64)
assert set(result.keys()) == {"chosen", "rejected", "chosen_mask", "rejected_mask"}
assert len(result["chosen"]) == len(result["chosen_mask"])
assert len(result["rejected"]) == len(result["rejected_mask"])
assert result["chosen_mask"][0] == 0
assert any(m == 1 for m in result["chosen_mask"])
assert result["rejected_mask"][0] == 0
def test_dpo_tokenize_malformed_record():
"""dpo_tokenize returns None for missing fields."""
class FakeTokenizer:
def apply_chat_template(
self, messages, tokenize=True, add_generation_prompt=True
):
return [1]
assert dpo_tokenize({}, FakeTokenizer()) is None
assert dpo_tokenize({"prompt": "a"}, FakeTokenizer()) is None
assert dpo_tokenize({"prompt": "a", "chosen": "b"}, FakeTokenizer()) is None
def test_dpo_jsonl_lazy_load(base_test_env):
"""DPODataset loads raw JSONL with tokenizer_path → lazy processor."""
test_dir = base_test_env["test_dir"]
tokenizer_path = _save_test_tokenizer(test_dir, base_test_env["tokenizer"])
records = [
{"input": "Hello", "chosen": "world", "rejected": "earth"},
{"input": "Foo", "chosen": "bar", "rejected": "baz"},
]
path = _write_dpo_jsonl(test_dir, records)
ds = DatasetFactory.load(
train_type="dpo",
load_path=path,
window_size=0,
tokenizer_path=tokenizer_path,
)
assert len(ds) == 2
assert ds.store.num_records == 2
assert ds.store._processor is not None
item = ds[0]
assert set(item.keys()) == {"chosen", "rejected", "chosen_mask", "rejected_mask"}
assert item["chosen"].dtype == torch.long
assert item["chosen_mask"].dtype == torch.bool
assert item["chosen"].shape == item["chosen_mask"].shape
assert item["chosen"].shape == item["rejected"].shape
def test_dpo_jsonl_lazy_no_tokenizer():
"""DPODataset on jsonl without tokenizer_path falls back to eager
(which requires dataset_config.json, so it should raise)."""
with tempfile.TemporaryDirectory() as d:
path = os.path.join(d, "dpo.jsonl")
with open(path, "w") as f:
f.write(json.dumps({"input": "a", "chosen": "b", "rejected": "c"}) + "\n")
with pytest.raises(FileNotFoundError, match="dataset_config.json"):
DatasetFactory.load(
train_type="dpo",
load_path=path,
window_size=0,
)
def test_jsonl_store_lazy_len_returns_record_count(base_test_env):
"""JsonlStore in lazy mode: len() returns record count, not tokens."""
test_dir = base_test_env["test_dir"]
records = [{"input": str(i), "chosen": "c", "rejected": "r"} for i in range(5)]
path = _write_dpo_jsonl(test_dir, records)
store = JsonlStore()
store.load(path, processor=lambda r: {"chosen": torch.tensor([1, 2])})
assert len(store) == 5
assert store.num_records == 5
def test_jsonl_store_eager_len_returns_token_count(base_test_env):
"""JsonlStore in eager mode: num_records reflects per-record count."""
test_dir = base_test_env["test_dir"]
tokenizer_path = _save_test_tokenizer(test_dir, base_test_env["tokenizer"])
data_dir = _write_jsonl_dataset(
test_dir,
tokenizer_path,
[{"text": "hello world"}, {"text": "foo bar"}],
config_overrides={
"preprocessing": {"max_seq_len": 128, "min_chars": 0},
"output": {"position_ids_mode": "none"},
},
)
store = JsonlStore()
store.load(data_dir)
assert store.num_records == 2
assert len(store.keys) > 0
def test_h5_store_dual_mode(base_test_env):
"""H5Store supports both fetch (stream) and fetch_record (record).
No window configured ``len(store)`` reflects the record count
(2). ``token_count`` retains the legacy stream length (128), and
token-stream access via :meth:`fetch` is still available for
callers that want explicit begin/end control.
"""
test_dir = base_test_env["test_dir"]
seq_length = 64
dummy_data = {
"chosen": [_rand_seq(seq_length), _rand_seq(seq_length)],
"rejected": [_rand_seq(seq_length), _rand_seq(seq_length)],
}
save_h5(test_dir, "dpo_data", dummy_data)
store = H5Store()
store.load(test_dir)
assert store.token_count == seq_length * 2
assert store.num_records == 2
assert len(store) == 2 # no window configured → record count
rec0 = store.fetch_record(0, "chosen")
assert rec0.shape == (seq_length,)
stream = store.fetch(0, 10, "chosen")
assert stream.shape == (10,)
# Window-configured view of the same data uses stream sample count:
# token_count=128, window_size=64 → num_samples = (128-1-64)//64 + 1 = 1
stream_view = H5Store(window_size=seq_length, stride=seq_length)
stream_view.load(test_dir)
assert len(stream_view) == 1
def test_mmap_store_stream_only_no_offsets(base_test_env):
"""MmapStore without offsets: num_records == 0, stream works.
No window configured ``len(store)`` is 0 (no iterate units).
``token_count`` remains 128 for raw token slicing, and ``fetch``
provides direct token-range access.
"""
test_dir = base_test_env["test_dir"]
seq_length = 128
dummy_data = {"sequence": [_rand_seq(seq_length)]}
save_bin(test_dir, dummy_data)
store = StoreFactory.create("bin")
store.load(test_dir)
assert store.token_count == seq_length
assert store.num_records == 0
assert len(store) == 0
chunk = store.fetch(0, 32, "sequence")
assert chunk.shape == (32,)
+52
View File
@@ -68,6 +68,28 @@ def test_chat_mask_only_assistant(chat_tokenizer, builder):
assert len(masked) > 0 assert len(masked) > 0
def test_chat_batch_matches_single(chat_tokenizer, builder):
config = make_chat_config()
items = [
{
"messages": [
{"role": "user", "content": "What is 2+2?"},
{"role": "assistant", "content": "4"},
]
},
{
"messages": [
{"role": "system", "content": "Be concise."},
{"role": "user", "content": "Say hello."},
{"role": "assistant", "content": "Hello."},
]
},
]
batch = builder.build_batch(items, config, chat_tokenizer)
single = [builder.build(item, config, chat_tokenizer) for item in items]
assert batch == single
@pytest.mark.parametrize( @pytest.mark.parametrize(
"mask_rules,mask_default,expect_nonzero", "mask_rules,mask_default,expect_nonzero",
[ [
@@ -152,6 +174,17 @@ def test_instruction_basic(test_tokenizer, builder):
assert len(result["sequence"]) == len(result["loss_mask"]) assert len(result["sequence"]) == len(result["loss_mask"])
def test_instruction_batch_matches_single(test_tokenizer, builder):
config = make_instruction_config()
items = [
{"prompt": "Translate to French: Hello", "response": "Bonjour"},
{"prompt": "Translate to German: Hello", "response": "Hallo"},
]
assert builder.build_batch(items, config, test_tokenizer) == [
builder.build(item, config, test_tokenizer) for item in items
]
def test_instruction_prompt_masked(test_tokenizer, builder): def test_instruction_prompt_masked(test_tokenizer, builder):
config = make_instruction_config() config = make_instruction_config()
item = {"prompt": "hello", "response": "world"} item = {"prompt": "hello", "response": "world"}
@@ -363,6 +396,25 @@ def test_grpo_basic(chat_tokenizer, builder):
assert result["rewards"] == [1.0, 0.5, 0.8, 0.2] assert result["rewards"] == [1.0, 0.5, 0.8, 0.2]
def test_grpo_batch_matches_single(chat_tokenizer, builder):
config = make_grpo_config()
items = [
{
"prompt": [{"role": "user", "content": "What is 2+2?"}],
"responses": ["4", "5"],
"rewards": [1.0, 0.0],
},
{
"prompt": [{"role": "user", "content": "Say hello."}],
"responses": ["Hello", "Hi"],
"rewards": [1.0, 0.5],
},
]
assert builder.build_batch(items, config, chat_tokenizer) == [
builder.build(item, config, chat_tokenizer) for item in items
]
def test_grpo_response_tokens_all_trained(chat_tokenizer, builder): def test_grpo_response_tokens_all_trained(chat_tokenizer, builder):
config = make_grpo_config() config = make_grpo_config()
item = { item = {
+6 -6
View File
@@ -1,4 +1,4 @@
from astrai.dataset import ResumableDistributedSampler from astrai.dataset import RDSampler
def test_random_sampler_consistency(random_dataset): def test_random_sampler_consistency(random_dataset):
@@ -6,8 +6,8 @@ def test_random_sampler_consistency(random_dataset):
dataset = random_dataset dataset = random_dataset
# Create two samplers with same seed # Create two samplers with same seed
sampler1 = ResumableDistributedSampler(dataset, seed=42) sampler1 = RDSampler(dataset, seed=42)
sampler2 = ResumableDistributedSampler(dataset, seed=42) sampler2 = RDSampler(dataset, seed=42)
indices1 = list(iter(sampler1)) indices1 = list(iter(sampler1))
indices2 = list(iter(sampler2)) indices2 = list(iter(sampler2))
@@ -20,8 +20,8 @@ def test_random_sampler_different_seeds(random_dataset):
dataset = random_dataset dataset = random_dataset
# Create two samplers with different seeds # Create two samplers with different seeds
sampler1 = ResumableDistributedSampler(dataset, seed=42) sampler1 = RDSampler(dataset, seed=42)
sampler2 = ResumableDistributedSampler(dataset, seed=123) sampler2 = RDSampler(dataset, seed=123)
indices1 = list(iter(sampler1)) indices1 = list(iter(sampler1))
indices2 = list(iter(sampler2)) indices2 = list(iter(sampler2))
@@ -35,7 +35,7 @@ def test_sampler_across_epochs(random_dataset):
dataset = random_dataset dataset = random_dataset
n = len(dataset) n = len(dataset)
sampler = ResumableDistributedSampler(dataset, seed=42) sampler = RDSampler(dataset, seed=42)
# Get indices for first epoch # Get indices for first epoch
epoch1_indices = list(iter(sampler)) epoch1_indices = list(iter(sampler))
+51
View File
@@ -231,3 +231,54 @@ def test_sample_with_frequency_penalty():
) )
assert tokens.shape == (1,) assert tokens.shape == (1,)
assert 0 <= tokens[0] < logits.size(-1) assert 0 <= tokens[0] < logits.size(-1)
def test_sample_return_logprobs_shape():
"""``return_logprobs=True`` returns ``[batch]`` logprobs aligned to tokens."""
logits = torch.tensor([[1.0, 2.0, 3.0], [3.0, 2.0, 1.0]])
out = sample(logits, temperature=1.0, return_logprobs=True)
tokens, logprobs = out
assert tokens.shape == (2,)
assert logprobs.shape == (2,)
def test_sample_return_logprobs_nonpositive():
"""Probabilities never exceed 1, so logprobs are always ≤ 0."""
torch.manual_seed(0)
logits = torch.randn(4, 50)
_, logprobs = sample(
logits, temperature=0.8, top_k=20, top_p=0.9, return_logprobs=True
)
assert torch.all(logprobs <= 1e-5)
def test_sample_return_logprobs_greedy_path():
"""Greedy decode (temperature 0) also returns logprobs."""
logits = torch.tensor([[1.0, 5.0, 2.0]])
tokens, logprobs = sample(logits, temperature=0.0, return_logprobs=True)
assert tokens[0].item() == 1
# log p(token=1) should equal log_softmax(logits)[1]
expected = torch.log_softmax(logits.float(), dim=-1)[0, 1]
assert torch.allclose(logprobs[0], expected, atol=1e-5)
def test_sample_return_logprobs_matches_manual_computation():
"""Returned logprob equals log_softmax(transformed_logits)[token]."""
torch.manual_seed(1)
logits = torch.randn(2, 30)
tokens, logprobs = sample(logits, temperature=0.7, top_p=0.95, return_logprobs=True)
# Recompute with the same pipeline
from astrai.inference.sample import (
SamplingPipeline,
TemperatureStrategy,
TopPStrategy,
)
pipeline = SamplingPipeline([TemperatureStrategy(0.7), TopPStrategy(0.95)])
transformed = pipeline.apply(logits.clone())
expected = torch.gather(
torch.log_softmax(transformed.float(), dim=-1),
-1,
tokens.unsqueeze(-1),
).squeeze(-1)
assert torch.allclose(logprobs, expected, atol=1e-5)
+126 -5
View File
@@ -14,11 +14,11 @@ def mock_model_and_tokenizer():
"""Create mock model and tokenizer.""" """Create mock model and tokenizer."""
mock_model = MagicMock() mock_model = MagicMock()
mock_model.config = MagicMock() mock_model.config = MagicMock()
mock_model.config.n_kv_heads = 8 mock_model.config.num_key_value_heads = 8
mock_model.config.n_heads = 8 mock_model.config.num_attention_heads = 8
mock_model.config.dim = 128 mock_model.config.hidden_size = 128
mock_model.config.n_layers = 2 mock_model.config.num_hidden_layers = 2
mock_model.config.max_len = 100 mock_model.config.max_position_embeddings = 100
mock_model.parameters.return_value = iter( mock_model.parameters.return_value = iter(
[MagicMock(dtype=torch.float32, device=torch.device("cpu"))] [MagicMock(dtype=torch.float32, device=torch.device("cpu"))]
) )
@@ -191,3 +191,124 @@ def test_prefill_skips_fully_cached_tasks(mock_model_and_tokenizer):
task_id = scheduler.add_task("short prompt", stream_callback=lambda t: None) task_id = scheduler.add_task("short prompt", stream_callback=lambda t: None)
scheduler.stop() scheduler.stop()
assert task_id.startswith("task_") assert task_id.startswith("task_")
def _make_real_scheduler(device):
"""Build a scheduler backed by a tiny real model for run_batch tests."""
from astrai.config.model_config import AutoRegressiveLMConfig
from astrai.model.transformer import AutoRegressiveLM
class _Tok:
stop_ids = [2]
def encode(self, texts, **_):
if isinstance(texts, str):
texts = [texts]
return [[b for b in t.encode("utf-8")] for t in texts]
def decode(self, ids, skip_special_tokens=True):
return bytes(b for b in ids if b > 2 or not skip_special_tokens).decode(
"utf-8", errors="ignore"
)
cfg = AutoRegressiveLMConfig(
vocab_size=200,
hidden_size=16,
num_attention_heads=2,
num_key_value_heads=1,
intermediate_size=32,
max_position_embeddings=64,
num_hidden_layers=2,
rms_norm_eps=1e-5,
)
model = AutoRegressiveLM(cfg).to(device=device).eval()
tokenizer = _Tok()
scheduler = InferenceScheduler(
model=model,
tokenizer=tokenizer,
max_batch_size=8,
max_seq_len=64,
max_prompt_len=64,
)
return scheduler, tokenizer, model
def test_run_batch_returns_token_sequences():
device = "cuda" if torch.cuda.is_available() else "cpu"
scheduler, _tok, _model = _make_real_scheduler(device)
try:
prompts = [[10, 20, 30], [5, 6, 7, 8]]
results = scheduler.run_batch(prompts, max_tokens=4, temperature=1.0)
assert len(results) == 2
for ids in results:
assert isinstance(ids, list)
assert len(ids) <= 4
assert all(0 <= i < 200 for i in ids)
finally:
scheduler.stop()
def test_run_batch_return_logprobs_aligned():
"""return_logprobs=True gives (token_ids, logprobs) tuples with equal len."""
device = "cuda" if torch.cuda.is_available() else "cpu"
scheduler, _tok, _model = _make_real_scheduler(device)
try:
prompts = [[10, 20, 30, 40]]
results = scheduler.run_batch(
prompts, max_tokens=5, temperature=1.0, return_logprobs=True
)
assert len(results) == 1
token_ids, logprobs = results[0]
assert len(token_ids) == len(logprobs)
assert all(lp <= 1e-5 for lp in logprobs) # logprobs ≤ 0
finally:
scheduler.stop()
def test_run_batch_respects_max_tokens():
device = "cuda" if torch.cuda.is_available() else "cpu"
scheduler, _tok, _model = _make_real_scheduler(device)
try:
prompts = [[10, 20, 30]]
results = scheduler.run_batch(prompts, max_tokens=3, temperature=1.0)
assert len(results[0]) <= 3
finally:
scheduler.stop()
def test_run_batch_stop_id_terminates():
"""A token matching stop_ids terminates generation for that prompt."""
device = "cuda" if torch.cuda.is_available() else "cpu"
scheduler, _tok, _model = _make_real_scheduler(device)
try:
prompts = [[10, 20, 30]]
results = scheduler.run_batch(prompts, max_tokens=32, temperature=1.0)
# If stop token 2 was produced, it is the last token
if results[0] and results[0][-1] == 2:
# No tokens after stop should exist (since we terminate)
assert 2 not in results[0][:-1]
finally:
scheduler.stop()
def test_run_batch_empty_prompts():
"""Empty prompt list yields empty result list."""
device = "cuda" if torch.cuda.is_available() else "cpu"
scheduler, _tok, _model = _make_real_scheduler(device)
try:
assert scheduler.run_batch([], max_tokens=4) == []
finally:
scheduler.stop()
def test_run_batch_too_long_prompt_skipped():
"""A prompt longer than max_seq_len yields an empty result slot."""
device = "cuda" if torch.cuda.is_available() else "cpu"
scheduler, _tok, _model = _make_real_scheduler(device)
try:
long = list(range(100)) # > max_seq_len=64
results = scheduler.run_batch([long, [10, 20]], max_tokens=2)
assert results[0] == []
assert len(results[1]) <= 2
finally:
scheduler.stop()
+10 -10
View File
@@ -12,13 +12,13 @@ from astrai.model.encoder import EmbeddingEncoder
TINY_CONFIG = dict( TINY_CONFIG = dict(
vocab_size=128, vocab_size=128,
dim=8, hidden_size=8,
n_heads=2, num_attention_heads=2,
n_kv_heads=1, num_key_value_heads=1,
dim_ffn=16, intermediate_size=16,
max_len=64, max_position_embeddings=64,
n_layers=2, num_hidden_layers=2,
norm_eps=1e-5, rms_norm_eps=1e-5,
) )
_device = "cuda" if torch.cuda.is_available() else "cpu" _device = "cuda" if torch.cuda.is_available() else "cpu"
@@ -42,7 +42,7 @@ def test_encoder_forward_pooling(pooling_type):
with torch.no_grad(): with torch.no_grad():
output = model(input_ids) output = model(input_ids)
assert output.shape == (batch_size, TINY_CONFIG["dim"]) assert output.shape == (batch_size, TINY_CONFIG["hidden_size"])
assert not torch.isnan(output).any() assert not torch.isnan(output).any()
@@ -60,7 +60,7 @@ def test_encoder_forward_with_padding():
with torch.no_grad(): with torch.no_grad():
output = model(input_ids, input_mask=input_mask) output = model(input_ids, input_mask=input_mask)
assert output.shape == (batch_size, TINY_CONFIG["dim"]) assert output.shape == (batch_size, TINY_CONFIG["hidden_size"])
assert not torch.isnan(output).any() assert not torch.isnan(output).any()
@@ -90,7 +90,7 @@ def test_encoder_from_transformer_checkpoint():
model = _make_model() model = _make_model()
state_dict = model.state_dict() state_dict = model.state_dict()
state_dict["lm_head.weight"] = torch.randn( state_dict["lm_head.weight"] = torch.randn(
TINY_CONFIG["vocab_size"], TINY_CONFIG["dim"], device=_device TINY_CONFIG["vocab_size"], TINY_CONFIG["hidden_size"], device=_device
) )
new_model = _make_model() new_model = _make_model()
+19 -10
View File
@@ -6,13 +6,13 @@ from astrai.model.transformer import AutoRegressiveLM
TINY_CONFIG = dict( TINY_CONFIG = dict(
vocab_size=128, vocab_size=128,
dim=8, hidden_size=8,
n_heads=2, num_attention_heads=2,
n_kv_heads=1, num_key_value_heads=1,
dim_ffn=16, intermediate_size=16,
max_len=64, max_position_embeddings=64,
n_layers=2, num_hidden_layers=2,
norm_eps=1e-5, rms_norm_eps=1e-5,
) )
@@ -58,8 +58,13 @@ CONFIGS = [
id="gqa_qk_norm", id="gqa_qk_norm",
), ),
pytest.param( pytest.param(
{**TINY_CONFIG, "attn_type": "gqa", "ffn_type": "mlp", "tie_weight": True}, {
id="gqa_tie_weight", **TINY_CONFIG,
"attn_type": "gqa",
"ffn_type": "mlp",
"tie_word_embeddings": True,
},
id="gqa_tie_word_embeddings",
), ),
] ]
@@ -82,7 +87,11 @@ def test_model_forward(config_kwargs):
assert "logits" in output assert "logits" in output
assert "hidden_states" in output assert "hidden_states" in output
assert output["logits"].shape == (batch_size, seq_len, config.vocab_size) assert output["logits"].shape == (batch_size, seq_len, config.vocab_size)
assert output["hidden_states"].shape == (batch_size, seq_len, config.dim) assert output["hidden_states"].shape == (
batch_size,
seq_len,
config.hidden_size,
)
assert not torch.isnan(output["logits"]).any() assert not torch.isnan(output["logits"]).any()
assert not torch.isnan(output["hidden_states"]).any() assert not torch.isnan(output["hidden_states"]).any()
+8 -8
View File
@@ -19,13 +19,13 @@ from astrai.model.components.lora import (
MODEL_KWARGS = dict( MODEL_KWARGS = dict(
vocab_size=1000, vocab_size=1000,
dim=64, hidden_size=64,
n_heads=4, num_attention_heads=4,
n_kv_heads=2, num_key_value_heads=2,
dim_ffn=128, intermediate_size=128,
n_layers=2, num_hidden_layers=2,
max_len=32, max_position_embeddings=32,
norm_eps=1e-5, rms_norm_eps=1e-5,
) )
@@ -192,7 +192,7 @@ def test_inject_lora_on_moe_model():
n_routed_experts=4, n_routed_experts=4,
n_shared_experts=1, n_shared_experts=1,
n_activated_experts=2, n_activated_experts=2,
dim_ffn=32, intermediate_size=32,
) )
inject_lora(model, r=4, alpha=8, target_modules={"up", "gate", "down"}) inject_lora(model, r=4, alpha=8, target_modules={"up", "gate", "down"})
assert _get_lora_count(model) > 0 assert _get_lora_count(model) > 0
+11 -11
View File
@@ -17,13 +17,13 @@ def transformer_test_env():
config = { config = {
"vocab_size": 1000, "vocab_size": 1000,
"dim": 8, "hidden_size": 8,
"n_heads": 2, "num_attention_heads": 2,
"n_kv_heads": 1, "num_key_value_heads": 1,
"dim_ffn": 16, "intermediate_size": 16,
"max_len": 64, "max_position_embeddings": 64,
"n_layers": 2, "num_hidden_layers": 2,
"norm_eps": 1e-5, "rms_norm_eps": 1e-5,
} }
with open(config_path, "w") as f: with open(config_path, "w") as f:
@@ -45,7 +45,7 @@ def test_tie_weight_init(transformer_test_env):
config_data = transformer_test_env["config"].copy() config_data = transformer_test_env["config"].copy()
# case 1: tie weight # case 1: tie weight
config_data["tie_weight"] = True config_data["tie_word_embeddings"] = True
with open(config_path, "w") as f: with open(config_path, "w") as f:
json.dump(config_data, f) json.dump(config_data, f)
@@ -63,7 +63,7 @@ def test_tie_weight_init(transformer_test_env):
assert not torch.equal(model.lm_head.weight, original_weight) assert not torch.equal(model.lm_head.weight, original_weight)
# case 2: not tie weight # case 2: not tie weight
config_data["tie_weight"] = False config_data["tie_word_embeddings"] = False
with open(config_path, "w") as f: with open(config_path, "w") as f:
json.dump(config_data, f) json.dump(config_data, f)
@@ -88,7 +88,7 @@ def test_model_save_load_with_tie_weight(transformer_test_env):
config_data = transformer_test_env["config"].copy() config_data = transformer_test_env["config"].copy()
# case 1: tie weight # case 1: tie weight
config_data["tie_weight"] = True config_data["tie_word_embeddings"] = True
config_path = os.path.join(test_dir, "config.json") config_path = os.path.join(test_dir, "config.json")
with open(config_path, "w") as f: with open(config_path, "w") as f:
@@ -108,7 +108,7 @@ def test_model_save_load_with_tie_weight(transformer_test_env):
assert "lm_head.weight" not in model.state_dict() assert "lm_head.weight" not in model.state_dict()
# case 2: not tie weight (form tie-weight state dict load) # case 2: not tie weight (form tie-weight state dict load)
config_data["tie_weight"] = False config_data["tie_word_embeddings"] = False
with open(config_path, "w") as f: with open(config_path, "w") as f:
json.dump(config_data, f) json.dump(config_data, f)
+8 -8
View File
@@ -13,16 +13,16 @@ class _FakeExecutor:
return model.state_dict() return model.state_dict()
def _make_config(vocab_size=200, max_len=64): def _make_config(vocab_size=200, max_position_embeddings=64):
return AutoRegressiveLMConfig( return AutoRegressiveLMConfig(
vocab_size=vocab_size, vocab_size=vocab_size,
dim=16, hidden_size=16,
n_heads=2, num_attention_heads=2,
n_kv_heads=1, num_key_value_heads=1,
dim_ffn=32, intermediate_size=32,
max_len=max_len, max_position_embeddings=max_position_embeddings,
n_layers=2, num_hidden_layers=2,
norm_eps=1e-5, rms_norm_eps=1e-5,
) )
+138
View File
@@ -0,0 +1,138 @@
"""End-to-end integration test for online DPO rollout."""
import os
from functools import partial
import pytest
import torch
from torch.utils.data import Dataset
from astrai.config import TrainConfig
from astrai.model.transformer import AutoRegressiveLM
from astrai.trainer.rollout import BaseRewardModel
from astrai.trainer.schedule import SchedulerFactory
from astrai.trainer.trainer import Trainer
_CHAT_TEMPLATE = (
"{% for message in messages %}"
"{% if message['role'] == 'system' %}"
"SYSTEM: {{ message['content'] }}\n"
"{% elif message['role'] == 'user' %}"
"USER: {{ message['content'] }}\n"
"{% elif message['role'] == 'assistant' %}"
"ASSISTANT: {{ message['content'] }}\n"
"{% endif %}"
"{% endfor %}"
"{% if add_generation_prompt %}ASSISTANT: {% endif %}"
)
class InstructionDataset(Dataset):
"""Toy instruction/input dataset for online RL rollout.
Each sample has an ``instruction`` and an optional ``input``; the
RolloutGenerator renders both through the tokenizer's chat template
so the prompt matches the SFT-trained format.
"""
_SAMPLES = [
{"instruction": "Hello", "input": ""},
{"instruction": "Tell me a story", "input": "about dragons"},
{"instruction": "Summarize", "input": "the article"},
{"instruction": "Translate", "input": "to French: hi"},
]
def __len__(self):
return len(self._SAMPLES)
def __getitem__(self, idx):
return dict(self._SAMPLES[idx])
class LengthRewardModel(BaseRewardModel):
"""Rewards each response by its (non-pad) token count.
Enough for DPO to distinguish chosen/rejected from the rollout group.
"""
def score(self, prompts, responses):
B = len(prompts)
G = len(responses[0]) if B else 0
rewards = torch.zeros(B, G)
for i in range(B):
for g in range(G):
rewards[i, g] = float(len(responses[i][g]))
return rewards
def instruction_collate_fn(batch):
"""Stack a list of instruction/input dicts into a batch dict of lists."""
return {
"instruction": [b["instruction"] for b in batch],
"input": [b.get("input", "") for b in batch],
}
def _model_fn(model_config):
return AutoRegressiveLM(model_config).to(dtype=torch.float32)
def _optimizer_fn(m):
return torch.optim.AdamW(m.parameters(), lr=1e-4)
def _scheduler_fn(optim):
return SchedulerFactory.create(
"cosine", optim, warmup_steps=1, lr_decay_steps=4, min_rate=0.05
)
@pytest.mark.integration
def test_online_dpo_end_to_end(base_test_env):
"""Run one epoch of online DPO with KV-cache-backed rollout."""
test_dir = base_test_env["test_dir"]
device = base_test_env["device"]
tokenizer = base_test_env["tokenizer"]
model_config = base_test_env["transformer_config"]
# Equip tokenizer with a chat template so RolloutGenerator can
# render instruction/input via apply_chat_template.
tokenizer.set_chat_template(_CHAT_TEMPLATE)
tokenizer.save_pretrained(test_dir)
model_fn = partial(_model_fn, model_config)
optimizer_fn = _optimizer_fn
scheduler_fn = _scheduler_fn
dataset = InstructionDataset()
train_config = TrainConfig(
strategy="online_dpo",
model_fn=model_fn,
dataset=dataset,
optimizer_fn=optimizer_fn,
scheduler_fn=scheduler_fn,
ckpt_dir=os.path.join(test_dir, "ckpt"),
log_dir=os.path.join(test_dir, "logs"),
n_epoch=1,
batch_per_device=2,
ckpt_interval=100,
grad_accum_steps=1,
random_seed=42,
device_type=device,
nprocs=1,
parallel_mode="none",
extra_kwargs={"beta": 0.1, "group_size": 2},
rollout_interval=1,
rollout_temperature=1.0,
rollout_top_k=0,
rollout_top_p=1.0,
rollout_max_tokens=4,
reward_model_fn=LengthRewardModel,
collate_fn=instruction_collate_fn,
)
trainer = Trainer(train_config)
trainer.train(param_path=test_dir)
assert os.path.isdir(os.path.join(test_dir, "ckpt"))
+355
View File
@@ -0,0 +1,355 @@
"""Unit tests for online rollout integration in :class:`BaseStrategy`.
Covers the shared rollout-trigger logic in ``BaseStrategy.__call__``
(runner injection, cache-driven refresh hook, ``step()`` callback) and
the per-strategy ``prepare_from_rollout`` mappings for both
:class:`GRPOStrategy` and :class:`DPOStrategy`.
"""
import pytest
import torch
from astrai.config.model_config import AutoRegressiveLMConfig
from astrai.model.transformer import AutoRegressiveLM
from astrai.trainer.rollout import RolloutResult
from astrai.trainer.strategy import (
DPOStrategy,
GRPOStrategy,
StrategyFactory,
)
class _FakeExecutor:
"""Executor stub tracking ``sync_gradients`` and providing unwrap_model."""
def __init__(self, sync_gradients=True):
self._sync_gradients = sync_gradients
@property
def sync_gradients(self):
return self._sync_gradients
def unwrap_model(self, model):
return model.state_dict()
def _make_config(vocab_size=200, max_position_embeddings=64):
return AutoRegressiveLMConfig(
vocab_size=vocab_size,
hidden_size=16,
num_attention_heads=2,
num_key_value_heads=1,
intermediate_size=32,
max_position_embeddings=max_position_embeddings,
num_hidden_layers=2,
rms_norm_eps=1e-5,
)
def _make_model(device):
cfg = _make_config()
return AutoRegressiveLM(cfg).to(device=device), cfg
def _make_frozen(model, device):
cfg = _make_config()
copy = AutoRegressiveLM(cfg).to(device=device)
copy.load_state_dict(model.state_dict())
copy.requires_grad_(False)
copy.eval()
return copy
def _make_rollout_result(B=2, G=4, P=6, R=8, device="cpu"):
return RolloutResult(
prompts=torch.randint(3, 200, (B, P), device=device),
prompt_mask=torch.ones(B, P, dtype=torch.bool, device=device),
responses=torch.randint(3, 200, (B, G, R), device=device),
response_mask=torch.ones(B, G, R, dtype=torch.bool, device=device),
rewards=torch.randn(B, G, device=device),
logprobs_old=torch.zeros(B, G, R, device=device),
)
class _RecordingRunner:
"""Fake RolloutRunner returning a fixed result with freshness tracking.
Freshness is ``True`` on the first call after construction or after
:meth:`swap_result`; ``False`` on subsequent cached calls mirroring
the real ``RolloutRunner`` contract without invoking generation.
"""
def __init__(self, result):
self.result = result
self.calls = 0
self.step_calls = 0
self._fresh = True
def __call__(self, batch):
self.calls += 1
fresh = self._fresh
self._fresh = False
return self.result, fresh
def step(self):
self.step_calls += 1
def swap_result(self, result):
self.result = result
self._fresh = True
@pytest.fixture
def device():
return "cuda" if torch.cuda.is_available() else "cpu"
def _make_grpo(device, executor=None):
model, _ = _make_model(device)
old_model = _make_frozen(model, device)
ref_model = _make_frozen(model, device)
return GRPOStrategy(
model=model,
device=device,
old_model=old_model,
ref_model=ref_model,
clip_eps=0.2,
kl_coef=0.01,
group_size=4,
model_fn=lambda c=_make_config(): AutoRegressiveLM(c).to(device=device),
executor=executor or _FakeExecutor(),
)
def _make_dpo(device, executor=None):
model, _ = _make_model(device)
ref_model = _make_frozen(model, device)
return DPOStrategy(
model=model,
device=device,
ref_model=ref_model,
beta=0.1,
reduction="sum",
model_fn=lambda c=_make_config(): AutoRegressiveLM(c).to(device=device),
executor=executor or _FakeExecutor(),
)
def test_factory_registers_online_aliases():
assert StrategyFactory.is_registered("online_grpo")
assert StrategyFactory.is_registered("online_dpo")
assert StrategyFactory._entries["online_grpo"] is GRPOStrategy
assert StrategyFactory._entries["online_dpo"] is DPOStrategy
def test_grpo_supports_online(device):
assert _make_grpo(device).supports_online() is True
def test_dpo_supports_online(device):
assert _make_dpo(device).supports_online() is True
def test_base_strategy_prepare_from_rollout_raises_by_default(device):
from astrai.trainer.strategy import BaseStrategy
class _Offline(BaseStrategy):
def compute_loss(self, batch):
return torch.tensor(0.0)
strat = _Offline(model=torch.nn.Linear(1, 1), device="cpu")
with pytest.raises(NotImplementedError):
strat.prepare_from_rollout(_make_rollout_result(device="cpu"))
def test_base_strategy_supports_online_default_false():
from astrai.trainer.strategy import BaseStrategy
class _Offline(BaseStrategy):
def compute_loss(self, batch):
return torch.tensor(0.0)
strat = _Offline(model=torch.nn.Linear(1, 1), device="cpu")
assert strat.supports_online() is False
def test_grpo_prepare_from_rollout_mapping(device):
strat = _make_grpo(device)
r = _make_rollout_result(device=device)
batch = strat.prepare_from_rollout(r)
assert batch["prompts"] is r.prompts
assert batch["prompt_mask"] is r.prompt_mask
assert batch["responses"] is r.responses
assert batch["masks"] is r.response_mask
assert batch["rewards"] is r.rewards
def test_dpo_prepare_from_rollout_picks_best_worst(device):
strat = _make_dpo(device)
r = _make_rollout_result(B=3, G=4, R=5, device=device)
batch = strat.prepare_from_rollout(r)
assert batch["chosen"].shape == (3, 5)
assert batch["rejected"].shape == (3, 5)
assert batch["chosen_mask"].shape == (3, 5)
assert batch["rejected_mask"].shape == (3, 5)
idx = torch.arange(3, device=device)
expected_best = r.responses[idx, r.rewards.argmax(dim=-1)]
expected_worst = r.responses[idx, r.rewards.argmin(dim=-1)]
assert torch.equal(batch["chosen"], expected_best)
assert torch.equal(batch["rejected"], expected_worst)
def test_call_without_runner_falls_back_to_compute_loss_grpo(device):
strat = _make_grpo(device)
batch = {
"prompts": torch.randint(3, 200, (2, 4), device=device),
"responses": torch.randint(3, 200, (2, 4, 6), device=device),
"masks": torch.ones(2, 4, 6, device=device),
"rewards": torch.randn(2, 4, device=device),
}
loss = strat(batch)
assert torch.isfinite(loss).item()
def test_call_with_runner_returns_finite_loss_grpo(device):
strat = _make_grpo(device)
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
assert torch.isfinite(loss).item()
def test_call_with_runner_returns_finite_loss_dpo(device):
strat = _make_dpo(device)
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
assert torch.isfinite(loss).item()
def test_call_invokes_runner_each_time(device):
strat = _make_grpo(device)
runner = _RecordingRunner(_make_rollout_result(device=device))
strat.set_rollout_runner(runner)
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
assert runner.calls == 2
def test_grpo_syncs_old_model_on_first_rollout(device):
strat = _make_grpo(device)
runner = _RecordingRunner(_make_rollout_result(device=device))
strat.set_rollout_runner(runner)
with torch.no_grad():
for p in strat.model.parameters():
p.add_(0.1)
old_before = {k: v.clone() for k, v in strat.old_model.state_dict().items()}
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
old_after = strat.old_model.state_dict()
synced = any(
not torch.allclose(old_before[k], old_after[k])
for k in old_before
if k in old_after
)
assert synced
def test_grpo_no_resync_when_same_cached_result(device):
strat = _make_grpo(device)
runner = _RecordingRunner(_make_rollout_result(device=device))
strat.set_rollout_runner(runner)
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
assert runner.calls == 2
assert runner.step_calls == 2
def test_grpo_resync_when_new_rollout_result(device):
strat = _make_grpo(device)
runner = _RecordingRunner(_make_rollout_result(device=device))
strat.set_rollout_runner(runner)
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
runner.swap_result(_make_rollout_result(device=device))
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
assert runner.calls == 2
assert runner.step_calls == 2
def test_dpo_no_sync_hook_when_new_rollout_result(device):
"""DPO has no old_model, so ``_on_rollout_refresh`` must be a no-op.
We verify by ensuring no AttributeError is raised (DPO has no
old_model) and that step is still called.
"""
strat = _make_dpo(device)
runner = _RecordingRunner(_make_rollout_result(device=device))
strat.set_rollout_runner(runner)
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
runner.swap_result(_make_rollout_result(device=device))
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
assert runner.step_calls == 2
def test_step_not_called_when_sync_gradients_false(device):
executor = _FakeExecutor(sync_gradients=False)
strat = _make_grpo(device, executor=executor)
runner = _RecordingRunner(_make_rollout_result(device=device))
strat.set_rollout_runner(runner)
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
assert runner.step_calls == 0
def test_step_called_when_sync_gradients_true(device):
executor = _FakeExecutor(sync_gradients=True)
strat = _make_grpo(device, executor=executor)
runner = _RecordingRunner(_make_rollout_result(device=device))
strat.set_rollout_runner(runner)
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
assert runner.step_calls == 1
def test_loss_is_differentiable_grpo(device):
strat = _make_grpo(device)
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
loss.backward()
has_grad = any(
p.grad is not None and p.grad.abs().sum() > 0 for p in strat.model.parameters()
)
assert has_grad
def test_loss_is_differentiable_dpo(device):
strat = _make_dpo(device)
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
loss.backward()
has_grad = any(
p.grad is not None and p.grad.abs().sum() > 0 for p in strat.model.parameters()
)
assert has_grad
def test_ref_and_old_model_not_updated_by_backward_grpo(device):
strat = _make_grpo(device)
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
loss.backward()
for p in strat.ref_model.parameters():
assert p.grad is None
for p in strat.old_model.parameters():
assert p.grad is None
def test_ref_model_not_updated_by_backward_dpo(device):
strat = _make_dpo(device)
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
loss.backward()
for p in strat.ref_model.parameters():
assert p.grad is None
+382
View File
@@ -0,0 +1,382 @@
"""Unit tests for the online rollout module."""
import pytest
import torch
from astrai.config.model_config import AutoRegressiveLMConfig
from astrai.inference.core.scheduler import InferenceScheduler
from astrai.model.transformer import AutoRegressiveLM
from astrai.trainer.rollout import (
BaseRewardModel,
RawRollout,
RolloutGenerator,
RolloutResult,
RolloutRunner,
)
_CHAT_TEMPLATE = (
"{% for message in messages %}"
"{% if message['role'] == 'system' %}SYSTEM: {{ message['content'] }}\n{% endif %}"
"{% if message['role'] == 'user' %}USER: {{ message['content'] }}\n{% endif %}"
"{% if message['role'] == 'assistant' %}ASSISTANT: {{ message['content'] }}\n{% endif %}"
"{% endfor %}"
"{% if add_generation_prompt %}ASSISTANT: {% endif %}"
)
class FakeTokenizer:
"""Minimal stub tokenizer with a chat template for rollout tests."""
stop_ids = [2]
def __init__(self):
from astrai.tokenize.chat_template import ChatTemplate
self._chat_template = ChatTemplate.from_string(_CHAT_TEMPLATE)
def encode(self, texts, **_):
if isinstance(texts, str):
texts = [texts]
return [[b for b in t.encode("utf-8")] for t in texts]
def decode(self, ids, skip_special_tokens=True):
if isinstance(ids, list):
return bytes(b for b in ids if b > 2).decode("utf-8", errors="ignore")
return str(ids)
def apply_chat_template(
self, messages, tokenize=True, add_generation_prompt=True, **_
):
rendered = self._chat_template.render(
messages=messages, add_generation_prompt=add_generation_prompt
)
if tokenize:
return (
self.encode(rendered)[0]
if isinstance(rendered, str)
else [self.encode(t)[0] for t in rendered]
)
return rendered
class ConstantRewardModel(BaseRewardModel):
"""Returns a constant reward for every response."""
def __init__(self, value: float = 1.0):
self.value = value
def score(self, prompts, responses):
B = len(prompts)
G = len(responses[0]) if B else 0
return torch.full((B, G), float(self.value))
class BadShapeRewardModel(BaseRewardModel):
def score(self, prompts, responses):
return torch.zeros(len(prompts))
class NonFiniteRewardModel(BaseRewardModel):
def score(self, prompts, responses):
B = len(prompts)
G = len(responses[0]) if B else 0
return torch.full((B, G), float("nan"))
def _make_config(vocab_size=200, max_position_embeddings=128):
return AutoRegressiveLMConfig(
vocab_size=vocab_size,
hidden_size=16,
num_attention_heads=2,
num_key_value_heads=1,
intermediate_size=32,
max_position_embeddings=max_position_embeddings,
num_hidden_layers=2,
rms_norm_eps=1e-5,
)
def _make_model(device):
cfg = _make_config()
m = AutoRegressiveLM(cfg).to(device=device)
m.eval()
return m, cfg
def _make_scheduler(model, tokenizer, max_batch_size=8, max_len=128):
return InferenceScheduler(
model=model,
tokenizer=tokenizer,
max_batch_size=max_batch_size,
max_seq_len=max_len,
max_prompt_len=max_len,
)
def _make_instruction_batch(n=2):
"""Build a batch of instruction+input prompts as lists of strings."""
instructions = [f"Tell me about topic {i}" for i in range(n)]
inputs = [f"context {i}" for i in range(n)]
return {"instruction": instructions, "input": inputs}
def test_raw_rollout_fields():
r = RawRollout(
prompts=torch.zeros(2, 4, dtype=torch.long),
prompt_mask=torch.ones(2, 4, dtype=torch.bool),
responses=torch.zeros(2, 3, 5, dtype=torch.long),
response_mask=torch.ones(2, 3, 5, dtype=torch.bool),
logprobs_old=torch.zeros(2, 3, 5),
)
assert r.prompts.shape == (2, 4)
assert r.responses.shape == (2, 3, 5)
assert r.prompt_texts == []
assert r.response_texts == []
def test_rollout_result_inherits_raw_rollout_fields():
r = RolloutResult(
prompts=torch.zeros(2, 4, dtype=torch.long),
prompt_mask=torch.ones(2, 4, dtype=torch.bool),
responses=torch.zeros(2, 3, 5, dtype=torch.long),
response_mask=torch.ones(2, 3, 5, dtype=torch.bool),
logprobs_old=torch.zeros(2, 3, 5),
rewards=torch.zeros(2, 3),
)
assert r.rewards.shape == (2, 3)
assert r.prompts.shape == (2, 4)
assert r.responses.shape == (2, 3, 5)
assert r.prompt_mask.shape == (2, 4)
def test_base_reward_model_is_abstract():
with pytest.raises(TypeError):
BaseRewardModel()
def test_constant_reward_model_shape():
rm = ConstantRewardModel(0.5)
out = rm.score(["a", "b"], [["x", "y", "z"], ["p", "q", "r"]])
assert out.shape == (2, 3)
assert torch.all(out == 0.5)
@pytest.fixture
def device():
return "cuda" if torch.cuda.is_available() else "cpu"
def _make_generator(device, **kw):
model, _ = _make_model(device)
tokenizer = FakeTokenizer()
scheduler = _make_scheduler(
model,
tokenizer,
max_batch_size=kw.get("max_batch_size", 8),
max_len=kw.get("max_position_embeddings", 128),
)
generator = RolloutGenerator(
scheduler=scheduler,
tokenizer=tokenizer,
max_tokens=kw.get("max_tokens", 8),
group_size=kw.get("group_size", 2),
temperature=kw.get("temperature", 1.0),
top_k=kw.get("top_k", 0),
top_p=kw.get("top_p", 1.0),
)
return generator, model
def test_rollout_generator_shapes(device):
gen, _ = _make_generator(device, group_size=3, max_tokens=5)
batch = _make_instruction_batch(n=2)
r = gen.generate(batch)
assert r.responses.shape == (2, 3, 5)
assert r.response_mask.shape == (2, 3, 5)
assert r.logprobs_old.shape == (2, 3, 5)
assert r.prompt_mask.shape == r.prompts.shape
assert len(r.prompt_texts) == 2
assert len(r.response_texts) == 2
assert len(r.response_texts[0]) == 3
def test_rollout_generator_uses_eval_and_restores_mode(device):
gen, model = _make_generator(device, group_size=1, max_tokens=2)
model.train()
seen_training = []
original = gen.scheduler.run_batch
def recording_run_batch(*args, **kwargs):
seen_training.append(model.training)
return original(*args, **kwargs)
gen.scheduler.run_batch = recording_run_batch
gen.generate(_make_instruction_batch(n=1))
assert seen_training == [False]
assert model.training is True
def test_rollout_generator_mask_matches_responses(device):
"""Positions beyond a response's length are pad (mask False)."""
gen, _ = _make_generator(device, group_size=2, max_tokens=6)
batch = _make_instruction_batch(n=2)
r = gen.generate(batch)
for i in range(2):
for g in range(2):
real = r.response_mask[i, g].sum().item()
assert r.responses[i, g, real:].sum() == 0
if real < r.logprobs_old.size(-1):
assert torch.all(r.logprobs_old[i, g, real:] == 0)
def test_rollout_generator_logprobs_are_nonpositive(device):
"""Behaviour-policy logprobs of sampled tokens should be ≤ 0."""
gen, _ = _make_generator(device, group_size=2, max_tokens=4)
batch = _make_instruction_batch(n=1)
r = gen.generate(batch)
for i in range(1):
for g in range(2):
mask = r.response_mask[i, g]
lp = r.logprobs_old[i, g][mask]
assert torch.all(lp <= 1e-5)
def test_rollout_generator_instruction_role_mapping(device):
"""instruction → system, input → user, output → assistant."""
gen, _ = _make_generator(device, group_size=1, max_tokens=4)
batch = {
"instruction": ["Be helpful"],
"input": ["What is 2+2?"],
"output": ["Four"],
}
r = gen.generate(batch)
text = r.prompt_texts[0]
assert "SYSTEM: Be helpful" in text
assert "USER: What is 2+2?" in text
assert "ASSISTANT: Four" in text
def test_rollout_generator_messages_format(device):
"""Rollout also accepts pre-built messages."""
gen, _ = _make_generator(device, group_size=2, max_tokens=4)
batch = {
"messages": [
[{"role": "user", "content": "Hello"}],
[{"role": "user", "content": "Goodbye"}],
]
}
r = gen.generate(batch)
assert r.responses.shape[0] == 2
assert len(r.prompt_texts) == 2
assert "Hello" in r.prompt_texts[0] or "USER" in r.prompt_texts[0]
def test_rollout_generator_bad_batch_raises(device):
"""Batch without messages or instruction raises a clear error."""
gen, _ = _make_generator(device)
with pytest.raises(
ValueError, match="must contain either 'messages' or 'instruction'"
):
gen.generate({"input_ids": torch.zeros(2, 4, dtype=torch.long)})
def _make_runner(device, **kw):
generator, model = _make_generator(
device,
group_size=kw.get("group_size", 2),
max_tokens=kw.get("max_tokens", 8),
max_batch_size=kw.get("max_batch_size", 8),
max_len=kw.get("max_position_embeddings", 128),
)
rm = ConstantRewardModel(1.0)
return (
RolloutRunner(
generator=generator,
reward_model=rm,
rollout_interval=kw.get("rollout_interval", 2),
),
model,
)
def test_rollout_runner_shapes(device):
runner, _ = _make_runner(device, group_size=3, max_tokens=5)
batch = _make_instruction_batch(n=2)
r, is_fresh = runner(batch)
assert is_fresh
assert r.responses.shape == (2, 3, 5)
assert r.response_mask.shape == (2, 3, 5)
assert r.rewards.shape == (2, 3)
assert r.logprobs_old.shape == (2, 3, 5)
assert len(r.prompt_texts) == 2
assert len(r.response_texts) == 2
assert len(r.response_texts[0]) == 3
def test_rollout_runner_cache_returns_stale_flag(device):
runner, _ = _make_runner(device, rollout_interval=10)
batch = _make_instruction_batch()
r1, fresh1 = runner(batch)
r2, fresh2 = runner(batch)
assert r1 is r2
assert fresh1 is True
assert fresh2 is False
def test_rollout_runner_refreshes_for_different_batch(device):
runner, _ = _make_runner(device, rollout_interval=100)
r1, fresh1 = runner(_make_instruction_batch(n=1))
batch2 = {"instruction": ["Different prompt"], "input": [""]}
r2, fresh2 = runner(batch2)
assert fresh1 is True
assert fresh2 is True
assert r2 is not r1
@pytest.mark.parametrize("reward_model", [BadShapeRewardModel, NonFiniteRewardModel])
def test_rollout_runner_rejects_invalid_rewards(device, reward_model):
generator, _ = _make_generator(device, group_size=2, max_tokens=2)
runner = RolloutRunner(generator, reward_model(), rollout_interval=1)
with pytest.raises(ValueError):
runner(_make_instruction_batch(n=1))
def test_rollout_runner_step_triggers_new_rollout(device):
runner, _ = _make_runner(device, rollout_interval=2)
batch = _make_instruction_batch()
r1, fresh1 = runner(batch)
assert fresh1 is True
runner.step()
# interval=2 means trigger when _steps_since_rollout >= 2; 1 step not enough
r2, fresh2 = runner(batch)
assert r2 is r1
assert fresh2 is False
runner.step()
# Now _steps_since_rollout == 2 -> re-rollout
r3, fresh3 = runner(batch)
assert r3 is not r1
assert fresh3 is True
def test_rollout_runner_clear_cache_forces_rerun(device):
runner, _ = _make_runner(device, rollout_interval=100)
batch = _make_instruction_batch()
r1, _ = runner(batch)
runner.clear_cache()
r2, fresh2 = runner(batch)
assert r2 is not r1
assert fresh2 is True
def test_rollout_runner_step_resets_counter(device):
runner, _ = _make_runner(device, rollout_interval=1)
batch = _make_instruction_batch()
r1, _ = runner(batch)
runner.step()
r2, fresh2 = runner(batch)
assert r2 is not r1
assert fresh2 is True
# Counter reset after rollout; second call w/o step should be cached.
r3, fresh3 = runner(batch)
assert r3 is r2
assert fresh3 is False
+177
View File
@@ -0,0 +1,177 @@
import json
import multiprocessing as mp
import os
import signal
import time
import pytest
import torch
import torch.optim as optim
from torch.utils.data import Dataset
from astrai.config import TrainConfig
from astrai.config.model_config import AutoRegressiveLMConfig
from astrai.model.transformer import AutoRegressiveLM
from astrai.parallel.signal_handler import register_signal_handlers
from astrai.trainer import Trainer
from astrai.trainer.schedule import SchedulerFactory
from astrai.trainer.train_context import TrainContext
class _PicklableDataset(Dataset):
def __init__(self, length=200, max_length=64, vocab_size=1000):
self.length = length
self.max_length = max_length
self.vocab_size = vocab_size
def __len__(self):
return self.length
def __getitem__(self, idx):
return {
"input_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
"target_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
}
def _build_model():
config = AutoRegressiveLMConfig(
vocab_size=1000,
hidden_size=8,
num_attention_heads=2,
num_key_value_heads=1,
intermediate_size=16,
max_position_embeddings=64,
num_hidden_layers=2,
rms_norm_eps=1e-5,
)
device = "cuda" if torch.cuda.is_available() else "cpu"
return AutoRegressiveLM(config).to(device=device)
class _ReadyCallback:
def __init__(self, ready_file):
self._ready_file = ready_file
def on_train_begin(self, context):
with open(self._ready_file, "w") as f:
f.write("ready")
f.flush()
os.fsync(f.fileno())
def _inner_run(batch_per_device, ckpt_interval, ckpt_dir, log_dir, ready_file):
dataset = _PicklableDataset()
def model_fn():
return _build_model()
def optimizer_fn(m):
return optim.AdamW(m.parameters(), lr=0.001)
def scheduler_fn(optim):
return SchedulerFactory.create(
"cosine", optim, warmup_steps=10, lr_decay_steps=10, min_rate=0.05
)
train_config = TrainConfig(
strategy="seq",
model_fn=model_fn,
dataset=dataset,
optimizer_fn=optimizer_fn,
scheduler_fn=scheduler_fn,
ckpt_dir=ckpt_dir,
log_dir=log_dir,
n_epoch=1,
batch_per_device=batch_per_device,
ckpt_interval=ckpt_interval,
grad_accum_steps=1,
random_seed=42,
device_type="cuda" if torch.cuda.is_available() else "cpu",
)
trainer = Trainer(train_config)
trainer.callbacks.insert(0, _ReadyCallback(ready_file))
trainer.train()
def _spawn_train_and_signal(ckpt_dir, sig, timeout=120):
log_dir = os.path.join(ckpt_dir, "logs")
ready_file = os.path.join(ckpt_dir, "ready.txt")
ctx = mp.get_context("spawn")
p = ctx.Process(
target=_inner_run,
args=(2, 1000, ckpt_dir, log_dir, ready_file),
)
p.start()
deadline = time.time() + 30
while time.time() < deadline:
if os.path.exists(ready_file):
with open(ready_file) as f:
if f.read().strip() == "ready":
break
if not p.is_alive():
break
time.sleep(0.5)
assert p.is_alive(), "Training process died before becoming ready"
os.kill(p.pid, sig)
p.join(timeout=timeout)
if p.is_alive():
p.kill()
p.join(timeout=5)
return p.exitcode
def test_context_stop_flag():
ctx = TrainContext()
assert not ctx.stop_requested
ctx.request_stop()
assert ctx.stop_requested
def test_register_signal_handlers():
ctx = TrainContext()
register_signal_handlers(ctx)
assert not ctx.stop_requested
os.kill(os.getpid(), signal.SIGTERM)
assert ctx.stop_requested
def test_sigterm_triggers_checkpoint_save(base_test_env):
exitcode = _spawn_train_and_signal(base_test_env["test_dir"], signal.SIGTERM)
assert exitcode == 0, f"Training process exited with code {exitcode} (expected 0)"
ckpt_dir = base_test_env["test_dir"]
meta_files = []
for root, dirs, files in os.walk(ckpt_dir):
for f in files:
if f == "meta.json":
meta_files.append(os.path.join(root, f))
assert len(meta_files) > 0, f"No checkpoint meta.json found in {ckpt_dir}"
with open(meta_files[-1]) as f:
meta = json.load(f)
assert "consumed_samples" in meta
assert meta["consumed_samples"] >= 0
@pytest.mark.slow
def test_sigint_triggers_checkpoint_save(base_test_env):
exitcode = _spawn_train_and_signal(base_test_env["test_dir"], signal.SIGINT)
assert exitcode == 0, f"Training process exited with code {exitcode} (expected 0)"
ckpt_dir = base_test_env["test_dir"]
meta_files = []
for root, dirs, files in os.walk(ckpt_dir):
for f in files:
if f == "meta.json":
meta_files.append(os.path.join(root, f))
assert len(meta_files) > 0, f"No checkpoint meta.json found in {ckpt_dir}"