34 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
78 changed files with 4341 additions and 1307 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
+239 -54
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
@@ -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()
@@ -166,13 +172,6 @@ classDiagram
+__getitem__(index) Dict +__getitem__(index) Dict
} }
class RecordDataset {
+Optional[Callable] processor
+load(load_path, storage_type)
+__getitem__(index)
+__len__()
}
class DPODataset { class DPODataset {
+__getitem__(index) Dict +__getitem__(index) Dict
} }
@@ -222,7 +221,12 @@ classDiagram
+fetch_record(index, keys) +fetch_record(index, keys)
} }
class ResumableDistributedSampler { class JsonlSource {
+Path path
+load() List[dict]
}
class RDSampler {
+int epoch +int epoch
+int iter +int iter
} }
@@ -385,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 {
@@ -497,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
} }
@@ -544,6 +632,32 @@ classDiagram
+sync_old_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 {
+get_lr() List[float] +get_lr() List[float]
+step() +step()
@@ -857,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
@@ -881,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
@@ -969,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]
} }
} }
@@ -994,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 {
@@ -1030,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
@@ -1052,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
@@ -1104,9 +1258,8 @@ classDiagram
TrainCallback <|-- MetricCallback TrainCallback <|-- MetricCallback
BaseDataset <|-- SEQDataset BaseDataset <|-- SEQDataset
BaseDataset <|-- SFTDataset BaseDataset <|-- SFTDataset
BaseDataset <|-- RecordDataset BaseDataset <|-- DPODataset
RecordDataset <|-- DPODataset BaseDataset <|-- GRPODataset
RecordDataset <|-- GRPODataset
Store <|-- H5Store Store <|-- H5Store
Store <|-- MmapStore Store <|-- MmapStore
Store <|-- JsonlStore Store <|-- JsonlStore
@@ -1119,6 +1272,7 @@ classDiagram
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
@@ -1142,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
@@ -1169,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
@@ -1191,6 +1366,9 @@ classDiagram
Pipeline o-- PipelineConfig Pipeline o-- PipelineConfig
Pipeline o-- BaseMaskBuilder Pipeline o-- BaseMaskBuilder
Pipeline o-- AutoTokenizer Pipeline o-- AutoTokenizer
Pipeline o-- PackingStrategy
Pipeline o-- PositionIdStrategy
Pipeline o-- StoreWriter
TokenizeTransform o-- AutoTokenizer TokenizeTransform o-- AutoTokenizer
TokenizeTransform o-- BaseMaskBuilder TokenizeTransform o-- BaseMaskBuilder
@@ -1198,6 +1376,9 @@ classDiagram
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
@@ -1216,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
@@ -1232,6 +1414,9 @@ 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
@@ -1253,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, SingleOutputMaskBuilder, MultiOutputMaskBuilder, Pipeline, TokenizeTransform, filter_by_length, PackingStrategy, PackingStrategyFactory, plan_bfd, PositionIdStrategy, PositionIdStrategyFactory, StoreWriter, StoreWriterFactory, core (shared helpers) | 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** | BaseDatasetRecordDatasetDPO/GRPODataset, SEQDataset, SFTDataset, Store, Streamable, Recordable, H5Store, MmapStore, JsonlStore, 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 |
@@ -1268,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 |
@@ -1277,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 |
@@ -1287,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`
@@ -1296,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-19 > Document Update Time: 2026-07-20
+1 -1
View File
@@ -85,7 +85,7 @@ All backends normalise tensors into `Store._data[Dict[str, List[Tensor]]]` + `St
``` ```
DatasetFactory.load(train_type, load_path, window_size, stride=None, DatasetFactory.load(train_type, load_path, window_size, stride=None,
storage_type=None, tokenizer_path=None, storage_type=None, tokenizer_path=None,
max_len=2048, store=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)
+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 (None disables) | None | | `--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-19 > Document Update Time: 2026-07-20
+20 -10
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)
@@ -146,6 +146,19 @@ Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`. External sync of `ol
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 |
@@ -162,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])
``` ```
@@ -181,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
@@ -222,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-19 > Document Update Time: 2026-07-20
+1 -1
View File
@@ -1,4 +1,4 @@
__version__ = "1.3.10" __version__ = "1.3.11"
__author__ = "ViperEkura" __author__ = "ViperEkura"
from astrai.config import ( from astrai.config import (
+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"
+27 -1
View File
@@ -38,7 +38,7 @@ class TrainConfig(BaseConfig):
default=1, metadata={"help": "Number of iterations between steps."} default=1, metadata={"help": "Number of iterations between steps."}
) )
max_grad_norm: Optional[float] = field( max_grad_norm: Optional[float] = field(
default=None, default=1.0,
metadata={"help": "Maximum gradient norm. None disables clipping."}, metadata={"help": "Maximum gradient norm. None disables clipping."},
) )
gradient_checkpointing_modules: List[str] = field( gradient_checkpointing_modules: List[str] = field(
@@ -138,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()."},
+15 -3
View File
@@ -190,7 +190,8 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
- rewards: [G] - rewards: [G]
Output: Output:
- prompts: [B, P_max] - prompts: [B, P_max], left-padded
- prompt_mask: [B, P_max]
- responses: [B, G, R_max] - responses: [B, G, R_max]
- masks: [B, G, R_max] - masks: [B, G, R_max]
- rewards: [B, G] - rewards: [B, G]
@@ -201,13 +202,15 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
R_max = max(r.size(0) for b in batch for r in b["responses"]) R_max = max(r.size(0) for b in batch for r in b["responses"])
prompts = torch.zeros(B, P_max, dtype=torch.long) prompts = torch.zeros(B, P_max, dtype=torch.long)
prompt_mask = torch.zeros(B, P_max, dtype=torch.bool)
responses = torch.zeros(B, G, R_max, dtype=torch.long) responses = torch.zeros(B, G, R_max, dtype=torch.long)
masks = torch.zeros(B, G, R_max, dtype=torch.bool) masks = torch.zeros(B, G, R_max, dtype=torch.bool)
rewards = torch.zeros(B, G, dtype=torch.float32) rewards = torch.zeros(B, G, dtype=torch.float32)
for i, b in enumerate(batch): for i, b in enumerate(batch):
p_len = b["prompts"].size(0) p_len = b["prompts"].size(0)
prompts[i, :p_len] = b["prompts"] prompts[i, -p_len:] = b["prompts"]
prompt_mask[i, -p_len:] = True
rewards[i, : b["rewards"].size(0)] = b["rewards"] rewards[i, : b["rewards"].size(0)] = b["rewards"]
for g in range(min(G, len(b["responses"]))): for g in range(min(G, len(b["responses"]))):
r_len = b["responses"][g].size(0) r_len = b["responses"][g].size(0)
@@ -217,6 +220,7 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
return { return {
"prompts": prompts, "prompts": prompts,
"prompt_mask": prompt_mask,
"responses": responses, "responses": responses,
"masks": masks, "masks": masks,
"rewards": rewards, "rewards": rewards,
@@ -346,7 +350,15 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
if processor is not None: if processor is not None:
store.load(load_path, processor=processor, **kwargs) store.load(load_path, processor=processor, **kwargs)
else: else:
store.load(load_path, **kwargs) load_kwargs = dict(kwargs)
if (
tokenizer_path is not None
and storage_type == "jsonl"
and train_type in ("seq", "sft")
and "tokenizer_path" not in load_kwargs
):
load_kwargs["tokenizer_path"] = tokenizer_path
store.load(load_path, **load_kwargs)
return cls.create(train_type, store=store) return cls.create(train_type, store=store)
+38 -10
View File
@@ -56,6 +56,7 @@ from typing import Callable, Dict, List, Optional, Tuple, Union
import torch import torch
from torch import Tensor from torch import Tensor
from astrai.config.preprocess_config import PipelineConfig
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.preprocessing.transform import TokenizeTransform from astrai.preprocessing.transform import TokenizeTransform
from astrai.serialization import ( from astrai.serialization import (
@@ -545,18 +546,29 @@ class JsonlSource:
@StoreFactory.register("jsonl") @StoreFactory.register("jsonl")
class JsonlStore(Store, Streamable, Recordable): class JsonlStore(Store, Streamable, Recordable):
"""JSONL reader with two tokenisation modes. """JSONL reader with eager/lazy tokenisation modes.
A JSONL dataset is a ``.jsonl`` file or a directory of ``*.jsonl`` A JSONL dataset is a ``.jsonl`` file or a directory of ``*.jsonl``
files plus (optionally) a ``dataset_config.json`` describing the files plus (optionally) a ``dataset_config.json`` describing the
tokenization pipeline. tokenization pipeline.
Two modes, selected at :meth:`load` time: Three ways to supply an eager transform (first match wins):
- **Eager** (default): applies a :class:`TokenizeTransform` to every - **Explicit** (``transform=``): caller-built
record at load time and registers per-key tensors via :class:`TokenizeTransform` applied eagerly.
``_normalize``. Both ``fetch`` (stream) and ``fetch_record`` - **Config file**: ``dataset_config.json`` alongside the ``*.jsonl``
(record) work. files — loaded via :meth:`TokenizeTransform.from_config_file`.
- **Default messages** (``tokenizer_path=`` given, no config file):
a built-in chatml config that tokenises the ``messages`` field,
masking every role except ``assistant`` (loss on assistant only).
Lets SFT/SEQ train straight from a chat-style JSONL directory
without a hand-written config.
Two tokenisation modes, selected at :meth:`load` time:
- **Eager** (default): applies the transform to every record at load
time and registers per-key tensors via ``_normalize``. Both
``fetch`` (stream) and ``fetch_record`` (record) work.
- **Lazy** (``processor=fn`` passed): keeps raw records and defers - **Lazy** (``processor=fn`` passed): keeps raw records and defers
tokenisation to ``fetch_record``. Only record access works — tokenisation to ``fetch_record``. Only record access works —
``len(store)`` returns ``num_records``; stream primitives raise. ``len(store)`` returns ``num_records``; stream primitives raise.
@@ -565,6 +577,16 @@ class JsonlStore(Store, Streamable, Recordable):
CONFIG_NAME = "dataset_config.json" CONFIG_NAME = "dataset_config.json"
segments_are_records = True segments_are_records = True
_DEFAULT_MESSAGES_CONFIG = {
"version": 1,
"input": {
"sections": [{"field": "messages", "action": "$role", "template": True}]
},
"mask": {"system": "mask", "user": "mask", "assistant": "train"},
"mask_default": "mask",
"output": {"position_ids_mode": "doc_reset"},
}
def __init__( def __init__(
self, self,
window_size: int = 0, window_size: int = 0,
@@ -587,14 +609,20 @@ class JsonlStore(Store, Streamable, Recordable):
if transform is None: if transform is None:
root = Path(path) root = Path(path)
config_path = root / self.CONFIG_NAME if root.is_dir() else None config_path = root / self.CONFIG_NAME if root.is_dir() else None
if config_path is None or not config_path.exists(): if config_path is not None and config_path.exists():
transform = TokenizeTransform.from_config_file(str(config_path))
else:
tokenizer_path = kwargs.get("tokenizer_path")
if not tokenizer_path:
raise FileNotFoundError( raise FileNotFoundError(
f"JSONL dataset config not found. Expected " f"JSONL dataset config not found. Expected "
f"{self.CONFIG_NAME} alongside *.jsonl files, pass an " f"{self.CONFIG_NAME} alongside *.jsonl files, pass an "
f"explicit transform, or pass processor= for lazy " f"explicit transform, pass processor= for lazy "
f"on-the-fly tokenisation." f"on-the-fly tokenisation, or pass tokenizer_path= to "
f"use the built-in messages config."
) )
transform = TokenizeTransform.from_config_file(str(config_path)) config = PipelineConfig.from_dict(self._DEFAULT_MESSAGES_CONFIG)
transform = TokenizeTransform(config, tokenizer_path)
transformed = transform.apply(records) transformed = transform.apply(records)
self._normalize(transformed) self._normalize(transformed)
+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
+4 -12
View File
@@ -435,24 +435,13 @@ class ContiguousCacheView(CacheView):
pos = self._write_positions pos = self._write_positions
self._cache.k[layer_id, indices, pos] = k.squeeze(1) self._cache.k[layer_id, indices, pos] = k.squeeze(1)
self._cache.v[layer_id, indices, pos] = v.squeeze(1) self._cache.v[layer_id, indices, pos] = v.squeeze(1)
for s, p in zip(indices.tolist(), pos.tolist()):
cur = self._cache._slot_len.get(s, 0)
if p + 1 > cur:
self._cache._slot_len[s] = p + 1
else: else:
start_pos = self._total_len - seq_len 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]
@@ -528,6 +517,9 @@ class ContiguousCache(KVCache):
) -> 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)
for slot in slots:
if total_len > self._slot_len.get(slot, 0):
self._slot_len[slot] = total_len
return ContiguousCacheView( return ContiguousCacheView(
self, batch_indices, total_len, write_positions=write_positions self, batch_indices, total_len, write_positions=write_positions
) )
+56 -17
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,32 +104,30 @@ 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),
input_mask=input_mask,
paged_cache=self.kv_cache.bind_tasks( paged_cache=self.kv_cache.bind_tasks(
task_ids, task_ids,
total_len, total_len,
@@ -116,6 +138,23 @@ class Executor:
) )
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,
+118 -6
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,
@@ -194,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
+1
View File
@@ -81,6 +81,7 @@ class Task:
self.status = TaskStatus.PENDING self.status = TaskStatus.PENDING
self.output_ids: List[int] = [] self.output_ids: List[int] = []
self.output_logprobs: List[float] = []
self.input_tokens: int = 0 self.input_tokens: int = 0
self.output_tokens: int = 0 self.output_tokens: int = 0
self.arrival_time = time.time() self.arrival_time = time.time()
+49 -17
View File
@@ -276,7 +276,8 @@ class SamplingPipeline(BaseSamplingStrategy):
filter_value: float = -float("inf"), filter_value: float = -float("inf"),
input_ids: Optional[Tensor] = None, input_ids: Optional[Tensor] = None,
input_mask: Optional[Tensor] = None, input_mask: Optional[Tensor] = None,
) -> Tensor: return_logprobs: bool = False,
):
"""Apply strategies then sample (softmax + multinomial). """Apply strategies then sample (softmax + multinomial).
Short-circuits to ``argmax`` when temperature is exactly 0 Short-circuits to ``argmax`` when temperature is exactly 0
@@ -286,21 +287,41 @@ class SamplingPipeline(BaseSamplingStrategy):
logits: Raw logits ``[batch, vocab_size]``. logits: Raw logits ``[batch, vocab_size]``.
input_ids: Previously generated token IDs ``[batch, seq_len]``. input_ids: Previously generated token IDs ``[batch, seq_len]``.
input_mask: Boolean mask for ``input_ids`` padding. input_mask: Boolean mask for ``input_ids`` padding.
return_logprobs: If ``True``, return ``(tokens, logprobs)``
where ``logprobs[i]`` is the log-probability of
``tokens[i]`` under the (post-strategy) sampling
distribution.
Returns: Returns:
Sampled token IDs ``[batch]``. Sampled token IDs ``[batch]``, or when ``return_logprobs``
is ``True`` a ``(token_ids, chosen_logprobs)`` tuple.
""" """
for s in self.strategies: if self._is_greedy_pipeline():
if isinstance(s, TemperatureStrategy) and self._is_greedy(s.temperature): tokens = logits.argmax(dim=-1)
return logits.argmax(dim=-1) if not return_logprobs:
break return tokens
log_probs = torch.log_softmax(logits.float(), dim=-1)
chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1)
return tokens, chosen
return torch.multinomial( transformed = self.apply(logits, filter_value, input_ids, input_mask)
torch.softmax( log_probs = torch.log_softmax(transformed.float(), dim=-1)
self.apply(logits, filter_value, input_ids, input_mask), dim=-1 tokens = torch.multinomial(
), torch.softmax(transformed, dim=-1), num_samples=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()
@@ -313,10 +334,11 @@ def sample(
input_ids: Optional[Tensor] = None, input_ids: Optional[Tensor] = None,
input_mask: Optional[Tensor] = None, input_mask: Optional[Tensor] = None,
filter_value: float = -float("inf"), filter_value: float = -float("inf"),
) -> Tensor: return_logprobs: bool = False,
):
"""Apply sampling strategies then sample (softmax + multinomial). """Apply sampling strategies then sample (softmax + multinomial).
Shortcut for ``SamplingPipeline(...).sample(logits)``. Shortcut for ``SamplingPipeline(...).sample(logits, return_logprobs=)``.
When **temperature** is exactly 0 (scalar or single-element tensor) When **temperature** is exactly 0 (scalar or single-element tensor)
the function short-circuits to ``argmax`` for deterministic decode. the function short-circuits to ``argmax`` for deterministic decode.
@@ -327,12 +349,16 @@ 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]``.
""" """
if SamplingPipeline._is_greedy(temperature):
return logits.argmax(dim=-1)
return SamplingPipeline( return SamplingPipeline(
[ [
TemperatureStrategy(temperature), TemperatureStrategy(temperature),
@@ -340,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",
] ]
+119 -23
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
@@ -148,14 +159,7 @@ class BaseExecutor:
def grad_accum_steps(self) -> int: def grad_accum_steps(self) -> int:
return self.gradient_state.num_steps return self.gradient_state.num_steps
def clip_grad_norm(self, model: nn.Module, max_norm: Optional[float]) -> float: def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
if max_norm is None:
total_norm = torch.norm(
torch.stack(
[p.grad.norm(2) for p in model.parameters() if p.grad is not None]
)
)
return total_norm.item()
total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm) total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)
if isinstance(total_norm, torch.Tensor): if isinstance(total_norm, torch.Tensor):
return total_norm.item() return total_norm.item()
@@ -289,9 +293,7 @@ class FSDPExecutor(BaseExecutor):
return model.no_sync() return model.no_sync()
return contextlib.nullcontext() return contextlib.nullcontext()
def clip_grad_norm(self, model: nn.Module, max_norm: Optional[float]) -> float: def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
if max_norm is None:
return super().clip_grad_norm(model, max_norm)
if isinstance(model, FSDP) and self.use_distributed: if isinstance(model, FSDP) and self.use_distributed:
total_norm = model.clip_grad_norm_(max_norm) total_norm = model.clip_grad_norm_(max_norm)
if isinstance(total_norm, torch.Tensor): if isinstance(total_norm, torch.Tensor):
@@ -309,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()
+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)
+39 -10
View File
@@ -23,7 +23,6 @@ import tqdm
from astrai.config.preprocess_config import PipelineConfig from astrai.config.preprocess_config import PipelineConfig
from astrai.preprocessing.core import ( from astrai.preprocessing.core import (
build_preprocessing_components, build_preprocessing_components,
iter_raw_records,
primary_ids, primary_ids,
) )
from astrai.preprocessing.packing import PackingStrategyFactory from astrai.preprocessing.packing import PackingStrategyFactory
@@ -81,6 +80,9 @@ class Pipeline:
def transform(self, item: dict) -> Optional[dict]: def transform(self, item: dict) -> Optional[dict]:
return self.mask_builder.build(item, self.config, self.tokenizer) return self.mask_builder.build(item, self.config, self.tokenizer)
def transform_batch(self, items: list[dict]) -> list[Optional[dict]]:
return self.mask_builder.build_batch(items, self.config, self.tokenizer)
def run(self): def run(self):
domains: dict = defaultdict(lambda: defaultdict(list)) domains: dict = defaultdict(lambda: defaultdict(list))
total_tokens = 0 total_tokens = 0
@@ -89,19 +91,31 @@ class Pipeline:
pp = self.config.preprocessing pp = self.config.preprocessing
for item in tqdm.tqdm( progress = tqdm.tqdm(desc="Tokenizing", unit="docs", mininterval=0.5)
self._iter_items(), desc="Tokenizing", unit="docs", mininterval=0.5 stop = False
): for items in self._iter_batches(pp.batch_size):
if pp.max_items and count >= pp.max_items: progress.update(len(items))
break
try: try:
result = self.transform(item) results = self.transform_batch(items)
except Exception: except Exception:
logger.warning( logger.warning(
"Failed to process item #%d, skipping", count + 1, exc_info=True "Failed to process batch, retrying records individually",
exc_info=True,
) )
continue results = []
for item in items:
try:
results.append(self.transform(item))
except Exception:
logger.warning(
"Failed to process item, skipping", exc_info=True
)
results.append(None)
for result in results:
if pp.max_items and count >= pp.max_items:
stop = True
break
if result is None: if result is None:
continue continue
@@ -122,6 +136,10 @@ class Pipeline:
self._flush(domains, shard_idx) self._flush(domains, shard_idx)
domains.clear() domains.clear()
total_tokens = 0 total_tokens = 0
if stop:
break
progress.close()
if total_tokens > 0: if total_tokens > 0:
self._flush(domains, shard_idx) self._flush(domains, shard_idx)
@@ -150,6 +168,17 @@ class Pipeline:
continue continue
yield json.loads(line) yield json.loads(line)
def _iter_batches(self, batch_size: int):
batch_size = max(1, batch_size)
batch = []
for item in self._iter_items():
batch.append(item)
if len(batch) >= batch_size:
yield batch
batch = []
if batch:
yield batch
def _flush(self, domains, shard_idx): def _flush(self, domains, shard_idx):
for domain, keys in domains.items(): for domain, keys in domains.items():
idx = shard_idx[domain] idx = shard_idx[domain]
+1 -1
View File
@@ -100,7 +100,7 @@ def load_bin(file_path: str) -> Dict[str, List[Tensor]]:
arr = np.memmap( arr = np.memmap(
os.path.join(file_path, f"{key}.bin"), os.path.join(file_path, f"{key}.bin"),
dtype=info["dtype"], dtype=info["dtype"],
mode="r", mode="c",
shape=tuple(info["shape"]), shape=tuple(info["shape"]),
) )
segments[key] = [torch.from_numpy(arr)] segments[key] = [torch.from_numpy(arr)]
+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",
] ]
+53 -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."""
@@ -227,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
+158 -17
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,
@@ -99,6 +113,7 @@ class BaseStrategy(ABC):
self.device = device self.device = device
self.executor = kwargs.pop("executor", None) self.executor = kwargs.pop("executor", 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:
@@ -112,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.
@@ -238,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]]
@@ -260,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):
@@ -314,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)
@@ -321,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]
@@ -371,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
+113 -51
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
@@ -8,11 +9,14 @@ from torch.utils.data import DataLoader, random_split
from astrai.config.train_config import TrainConfig from astrai.config.train_config import TrainConfig
from astrai.dataset import RDSampler from astrai.dataset import RDSampler
from astrai.inference.core.scheduler import InferenceScheduler
from astrai.model.components.lora import inject_lora from astrai.model.components.lora import inject_lora
from astrai.parallel.executor import BaseExecutor, ExecutorFactory from astrai.parallel.executor import BaseExecutor, ExecutorFactory
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
@@ -175,15 +196,6 @@ class TrainContextBuilder:
collate_fn=cfg.collate_fn, collate_fn=cfg.collate_fn,
) )
context.model, context.optimizer, context.dataloader, context.scheduler = (
executor.prepare(
model,
context.optimizer,
context.dataloader,
context.scheduler,
)
)
if context.checkpoint and context.checkpoint.extra: if context.checkpoint and context.checkpoint.extra:
extra = context.checkpoint.extra extra = context.checkpoint.extra
for name in ("optimizer", "scheduler"): for name in ("optimizer", "scheduler"):
@@ -194,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)
@@ -214,4 +235,45 @@ class TrainContextBuilder:
**strategy_kwargs, **strategy_kwargs,
) )
# Enable online rollout when the train_type is an ``online_*`` variant.
is_online = cfg.strategy.startswith("online_")
if is_online:
if not context.strategy.supports_online():
raise ValueError(
f"Strategy '{cfg.strategy}' does not support online rollout"
)
if cfg.reward_model_fn is None:
raise ValueError("reward_model_fn is required for online RL strategies")
tokenizer = AutoTokenizer.from_pretrained(self._param_path)
reward_model = cfg.reward_model_fn()
group_size = strategy_kwargs.get("group_size", 1)
rollout_batch_size = group_size * max(1, cfg.batch_per_device)
max_seq_len = getattr(context.model.config, "max_position_embeddings", None)
scheduler = InferenceScheduler(
model=context.model,
tokenizer=tokenizer,
max_batch_size=rollout_batch_size,
max_seq_len=max_seq_len,
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);
} }
+58 -73
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;
} }
+55 -64
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);
+56 -106
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"]
+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")
+1 -2
View File
@@ -185,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:]
@@ -215,7 +215,6 @@ def _permute_choices(item: dict, rng: random.Random) -> tuple[dict, str]:
positional bias (e.g. always picking B). positional bias (e.g. always picking B).
""" """
letters = ("A", "B", "C", "D") letters = ("A", "B", "C", "D")
contents = [item[k] for k in letters]
perm = list(letters) perm = list(letters)
rng.shuffle(perm) rng.shuffle(perm)
permuted = {"question": item["question"]} permuted = {"question": item["question"]}
+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(
+5 -2
View File
@@ -56,7 +56,7 @@ def processor(
print(f" {len(prompts)} prompts loaded\n") print(f" {len(prompts)} prompts loaded\n")
if max_tokens is None: if max_tokens is None:
max_tokens = model.config.max_len max_tokens = model.config.max_position_embeddings
chunk_size = max(1, batch_size) chunk_size = max(1, batch_size)
@@ -185,7 +185,10 @@ if __name__ == "__main__":
"--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( parser.add_argument(
"--cache_len", "--cache_len",
+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,
+72 -12
View File
@@ -1,7 +1,7 @@
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
@@ -12,6 +12,7 @@ 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(
@@ -148,7 +149,7 @@ def parse_args() -> argparse.Namespace:
parser.add_argument( parser.add_argument(
"--max_grad_norm", "--max_grad_norm",
type=float, type=float,
default=None, default=1.0,
help="Max gradient norm for clipping. None disables clipping.", help="Max gradient norm for clipping. None disables clipping.",
) )
parser.add_argument( parser.add_argument(
@@ -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(
@@ -510,6 +562,8 @@ def train(
collate_fn = dpo_collate_fn collate_fn = dpo_collate_fn
elif train_type == "grpo": elif train_type == "grpo":
collate_fn = grpo_collate_fn 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,
@@ -544,6 +598,12 @@ def train(
extra_kwargs=strategy_kwargs, extra_kwargs=strategy_kwargs,
neftune_alpha=neftune_alpha, neftune_alpha=neftune_alpha,
collate_fn=collate_fn, 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,
) )
+90 -3
View File
@@ -654,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")
@@ -738,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()
@@ -851,8 +937,9 @@ 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
+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 = {
+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}"