Compare commits
58
Commits
8999ca89b8
...
v1.3.9
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8f89c82d55 | ||
|
|
21871197d7 | ||
|
|
4c35d36146 | ||
|
|
9aca62c26c | ||
|
|
b5cdea98ad | ||
|
|
69fecaf387 | ||
|
|
fd6d25ad86 | ||
|
|
2c3cef1c87 | ||
|
|
89ece26c25 | ||
|
|
2c0b5d0b5e | ||
|
|
a4ae7d17fb | ||
|
|
8a8550184f | ||
|
|
b8b439b713 | ||
|
|
41cd40363a | ||
|
|
d923ebe38d | ||
|
|
29b0423c4e | ||
|
|
88f8dca2c2 | ||
|
|
9027fdc546 | ||
|
|
cbd140340d | ||
|
|
988e01314d | ||
|
|
7ba43a7c6f | ||
|
|
dea59f7e1d | ||
|
|
85dc771460 | ||
|
|
2c5629b81d | ||
|
|
841a582b28 | ||
|
|
c8567a6f65 | ||
|
|
8035be9b1f | ||
|
|
e9b03f4fca | ||
|
|
fd65b9bc23 | ||
|
|
9ebaea840f | ||
|
|
6adc221c10 | ||
|
|
9e63cb9ed0 | ||
|
|
4225518cf3 | ||
|
|
c50adbaac0 | ||
|
|
536dbc0c9a | ||
|
|
4af7acd449 | ||
|
|
53ed52b4b8 | ||
|
|
f1cc7cedce | ||
|
|
ddc4bd1cf6 | ||
|
|
cc36530c73 | ||
|
|
11fa807cfc | ||
|
|
bcdd93e0eb | ||
|
|
579b8c3129 | ||
|
|
d7da51569f | ||
|
|
e8e228d035 | ||
|
|
2579658e15 | ||
|
|
f0cd0134c6 | ||
|
|
abb96996f8 | ||
|
|
bbe6ff2d8f | ||
|
|
db9b39b084 | ||
|
|
849e1e00a3 | ||
|
|
5416c2e8fb | ||
|
|
599a51f4f7 | ||
|
|
17d6eaa2f2 | ||
|
|
2d908639e9 | ||
|
|
c7158418dd | ||
|
|
4d3c9341c1 | ||
|
|
4e508afa2d |
@@ -0,0 +1,71 @@
|
|||||||
|
name: Release
|
||||||
|
|
||||||
|
on:
|
||||||
|
push:
|
||||||
|
tags:
|
||||||
|
- "v*"
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
build-pure:
|
||||||
|
name: Build pure-Python wheel
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
steps:
|
||||||
|
- uses: actions/checkout@v4
|
||||||
|
- uses: actions/setup-python@v5
|
||||||
|
with:
|
||||||
|
python-version: "3.12"
|
||||||
|
|
||||||
|
- name: Build wheel (no CUDA)
|
||||||
|
run: |
|
||||||
|
pip wheel . --no-deps -w dist/
|
||||||
|
|
||||||
|
- uses: actions/upload-artifact@v4
|
||||||
|
with:
|
||||||
|
name: pure-wheel
|
||||||
|
path: dist/*.whl
|
||||||
|
|
||||||
|
build-cuda-linux:
|
||||||
|
name: Build CUDA wheel (Linux)
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
steps:
|
||||||
|
- uses: actions/checkout@v4
|
||||||
|
- uses: actions/setup-python@v5
|
||||||
|
with:
|
||||||
|
python-version: "3.12"
|
||||||
|
|
||||||
|
- name: Install torch (CUDA 12.8)
|
||||||
|
run: |
|
||||||
|
pip install torch --index-url https://download.pytorch.org/whl/cu128
|
||||||
|
|
||||||
|
- name: Setup CUDA
|
||||||
|
uses: Jimver/cuda-toolkit@v0.2.35
|
||||||
|
with:
|
||||||
|
cuda: "12.8.0"
|
||||||
|
|
||||||
|
- name: Build wheel (with CUDA kernels)
|
||||||
|
run: |
|
||||||
|
CSRC_KERNELS=true pip wheel . --no-deps --no-build-isolation -w dist/
|
||||||
|
|
||||||
|
- uses: actions/upload-artifact@v4
|
||||||
|
with:
|
||||||
|
name: cuda-wheel-linux
|
||||||
|
path: dist/*.whl
|
||||||
|
|
||||||
|
release:
|
||||||
|
name: Attach wheels to release
|
||||||
|
needs: [build-pure, build-cuda-linux]
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
permissions:
|
||||||
|
contents: write
|
||||||
|
steps:
|
||||||
|
- uses: actions/download-artifact@v4
|
||||||
|
with:
|
||||||
|
pattern: "*-wheel"
|
||||||
|
merge-multiple: true
|
||||||
|
|
||||||
|
- name: Create release & upload assets
|
||||||
|
uses: softprops/action-gh-release@v2
|
||||||
|
with:
|
||||||
|
files: ./*.whl
|
||||||
|
tag_name: ${{ github.ref_name }}
|
||||||
|
generate_release_notes: true
|
||||||
+12
-1
@@ -7,8 +7,14 @@
|
|||||||
# Allow specific file types and root files
|
# Allow specific file types and root files
|
||||||
!astrai/**/*.py
|
!astrai/**/*.py
|
||||||
!scripts/**/*.py
|
!scripts/**/*.py
|
||||||
!scripts/**/*.sh
|
|
||||||
!tests/**/*.py
|
!tests/**/*.py
|
||||||
|
!csrc/**/*.py
|
||||||
|
|
||||||
|
!csrc/**/*.cu
|
||||||
|
!csrc/**/*.h
|
||||||
|
!csrc/**/*.cuh
|
||||||
|
|
||||||
|
!scripts/**/*.sh
|
||||||
|
|
||||||
# Allow GitHub files
|
# Allow GitHub files
|
||||||
!/.github/**
|
!/.github/**
|
||||||
@@ -23,3 +29,8 @@
|
|||||||
!/LICENSE
|
!/LICENSE
|
||||||
!/pyproject.toml
|
!/pyproject.toml
|
||||||
!/README.md
|
!/README.md
|
||||||
|
# Allow extension modules (only source .py)
|
||||||
|
!/astrai/extension/**/*.py
|
||||||
|
|
||||||
|
# Allow build files
|
||||||
|
!/setup.py
|
||||||
|
|||||||
@@ -9,7 +9,7 @@
|
|||||||
<div align="center">
|
<div align="center">
|
||||||
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
|
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
|
||||||
<img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license">
|
<img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license">
|
||||||
<img src="https://img.shields.io/github/v/release/ViperEkura/AstrAI?label=Release&color=76bad9" alt="release">
|
<img src="https://img.shields.io/github/v/tag/ViperEkura/AstrAI?label=Release&color=76bad9" alt="release">
|
||||||
<img src="https://img.shields.io/github/stars/ViperEkura/AstrAI?style=flat&label=Stars&color=76bad9" alt="stars">
|
<img src="https://img.shields.io/github/stars/ViperEkura/AstrAI?style=flat&label=Stars&color=76bad9" alt="stars">
|
||||||
<img src="https://img.shields.io/github/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks">
|
<img src="https://img.shields.io/github/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks">
|
||||||
</div>
|
</div>
|
||||||
@@ -59,8 +59,9 @@ End-to-end walkthrough in 5 steps:
|
|||||||
```bash
|
```bash
|
||||||
git clone https://github.com/ViperEkura/AstrAI.git
|
git clone https://github.com/ViperEkura/AstrAI.git
|
||||||
cd AstrAI
|
cd AstrAI
|
||||||
pip install -e .
|
pip install -e . # pure PyTorch (no CUDA kernels)
|
||||||
# pip install -e ".[dev]" # optional: dev dependencies (pytest, ruff)
|
# CSRC_KERNELS=true pip install -e . --no-build-isolation # optional: fused CUDA kernels
|
||||||
|
# pip install -e ".[dev]" # dev dependencies (pytest, ruff)
|
||||||
```
|
```
|
||||||
|
|
||||||
**2. Download model**
|
**2. Download model**
|
||||||
@@ -102,9 +103,7 @@ nohup python scripts/tools/train.py \
|
|||||||
--warmup_ratio=0.05 \
|
--warmup_ratio=0.05 \
|
||||||
--max_lr=1e-4 \
|
--max_lr=1e-4 \
|
||||||
--max_grad_norm=1.0 \
|
--max_grad_norm=1.0 \
|
||||||
--adamw_beta1=0.9 \
|
--weight_decay=0.1 \
|
||||||
--adamw_beta2=0.95 \
|
|
||||||
--adamw_weight_decay=0.01 \
|
|
||||||
--window_size=2048 \
|
--window_size=2048 \
|
||||||
--ckpt_interval=10000 \
|
--ckpt_interval=10000 \
|
||||||
--ckpt_dir=./checkpoint \
|
--ckpt_dir=./checkpoint \
|
||||||
|
|||||||
@@ -15,7 +15,7 @@
|
|||||||
<div align="center">
|
<div align="center">
|
||||||
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
|
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
|
||||||
<img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license">
|
<img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license">
|
||||||
<img src="https://img.shields.io/github/v/release/ViperEkura/AstrAI?label=Release&color=76bad9" alt="release">
|
<img src="https://img.shields.io/github/v/tag/ViperEkura/AstrAI?label=Release&color=76bad9" alt="release">
|
||||||
<img src="https://img.shields.io/github/stars/ViperEkura/AstrAI?style=flat&label=Stars&color=76bad9" alt="stars">
|
<img src="https://img.shields.io/github/stars/ViperEkura/AstrAI?style=flat&label=Stars&color=76bad9" alt="stars">
|
||||||
<img src="https://img.shields.io/github/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks">
|
<img src="https://img.shields.io/github/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks">
|
||||||
</div>
|
</div>
|
||||||
@@ -65,8 +65,9 @@
|
|||||||
```bash
|
```bash
|
||||||
git clone https://github.com/ViperEkura/AstrAI.git
|
git clone https://github.com/ViperEkura/AstrAI.git
|
||||||
cd AstrAI
|
cd AstrAI
|
||||||
pip install -e .
|
pip install -e . # 纯 PyTorch(不含 CUDA 内核)
|
||||||
# pip install -e ".[dev]" # 可选:开发依赖(pytest, ruff)
|
# CSRC_KERNELS=true pip install -e . --no-build-isolation # 可选:融合 CUDA 内核加速
|
||||||
|
# pip install -e ".[dev]" # 可选:开发依赖(pytest, ruff)
|
||||||
```
|
```
|
||||||
|
|
||||||
**2. 下载模型**
|
**2. 下载模型**
|
||||||
@@ -108,9 +109,7 @@ nohup python scripts/tools/train.py \
|
|||||||
--warmup_ratio=0.05 \
|
--warmup_ratio=0.05 \
|
||||||
--max_lr=1e-4 \
|
--max_lr=1e-4 \
|
||||||
--max_grad_norm=1.0 \
|
--max_grad_norm=1.0 \
|
||||||
--adamw_beta1=0.9 \
|
--weight_decay=0.1 \
|
||||||
--adamw_beta2=0.95 \
|
|
||||||
--adamw_weight_decay=0.01 \
|
|
||||||
--window_size=2048 \
|
--window_size=2048 \
|
||||||
--ckpt_interval=10000 \
|
--ckpt_interval=10000 \
|
||||||
--ckpt_dir=./checkpoint \
|
--ckpt_dir=./checkpoint \
|
||||||
|
|||||||
+83
-49
@@ -63,7 +63,6 @@ classDiagram
|
|||||||
+Optional[int] n_heads
|
+Optional[int] n_heads
|
||||||
+Optional[int] n_kv_heads
|
+Optional[int] n_kv_heads
|
||||||
+Optional[bool] use_qk_norm
|
+Optional[bool] use_qk_norm
|
||||||
+Optional[bool] use_gated_attention
|
|
||||||
+str ffn_type
|
+str ffn_type
|
||||||
+Optional[dict] rope_scaling
|
+Optional[dict] rope_scaling
|
||||||
+Optional[str] pooling_type
|
+Optional[str] pooling_type
|
||||||
@@ -125,7 +124,6 @@ classDiagram
|
|||||||
+str ckpt_dir
|
+str ckpt_dir
|
||||||
+int ckpt_interval
|
+int ckpt_interval
|
||||||
+str log_dir
|
+str log_dir
|
||||||
+int log_interval
|
|
||||||
+List[str] metrics
|
+List[str] metrics
|
||||||
+Optional[LoRAConfig] lora
|
+Optional[LoRAConfig] lora
|
||||||
+int random_seed
|
+int random_seed
|
||||||
@@ -559,7 +557,7 @@ classDiagram
|
|||||||
}
|
}
|
||||||
|
|
||||||
class GradientCheckpointingCallback {
|
class GradientCheckpointingCallback {
|
||||||
+tuple modules
|
+Optional[List[type]] modules
|
||||||
+on_train_begin(context)
|
+on_train_begin(context)
|
||||||
+on_train_end(context)
|
+on_train_end(context)
|
||||||
}
|
}
|
||||||
@@ -573,31 +571,29 @@ classDiagram
|
|||||||
+on_batch_end(context)
|
+on_batch_end(context)
|
||||||
+on_train_end(context)
|
+on_train_end(context)
|
||||||
+on_error(context)
|
+on_error(context)
|
||||||
+save_extra(context) dict$
|
+save_extra(context) dict
|
||||||
}
|
}
|
||||||
|
|
||||||
class ProgressBarCallback {
|
class ProgressBarCallback {
|
||||||
+int num_epoch
|
+int num_epoch
|
||||||
+int log_interval
|
+int log_interval
|
||||||
+IO file
|
+IO file
|
||||||
|
+tqdm progress_bar
|
||||||
+on_epoch_begin(context)
|
+on_epoch_begin(context)
|
||||||
+on_batch_end(context)
|
+on_optimizer_step(context)
|
||||||
+on_epoch_end(context)
|
+on_epoch_end(context)
|
||||||
}
|
}
|
||||||
|
|
||||||
class MetricLoggerCallback {
|
class MetricCallback {
|
||||||
+Path log_dir
|
+Path log_dir
|
||||||
+int save_interval
|
+int save_interval
|
||||||
+int log_interval
|
|
||||||
+List[str] metrics
|
+List[str] metrics
|
||||||
+on_batch_end(context)
|
+int val_step
|
||||||
|
+on_optimizer_step(context)
|
||||||
|
+on_epoch_end(context)
|
||||||
+on_train_end(context)
|
+on_train_end(context)
|
||||||
+on_error(context)
|
+on_error(context)
|
||||||
}
|
|
||||||
|
|
||||||
class ValidationCallback {
|
|
||||||
-_run_validation(context)
|
-_run_validation(context)
|
||||||
+on_optimizer_step(context)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
class CallbackFactory {
|
class CallbackFactory {
|
||||||
@@ -684,20 +680,44 @@ classDiagram
|
|||||||
}
|
}
|
||||||
|
|
||||||
class KVCache {
|
class KVCache {
|
||||||
-PagePool _pool
|
<<abstract>>
|
||||||
-Storage _storage
|
|
||||||
-TaskTable _table
|
|
||||||
+int page_size
|
|
||||||
+task_alloc(task_id, prompt_ids) bool
|
+task_alloc(task_id, prompt_ids) bool
|
||||||
+task_free(task_id)
|
+task_free(task_id)
|
||||||
+task_extend(task_id, pos) bool
|
+task_extend(task_id, pos) bool
|
||||||
+task_cached(task_id) int
|
+task_cached(task_id) int
|
||||||
+task_record_hashes(task_id, prompt_ids, start_logical_page)
|
+task_record_hashes(task_id, prompt_ids, start_logical_page)
|
||||||
+make_table_tensor(task_ids, device) Tensor
|
+bind_tasks(task_ids, total_len, device) CacheView
|
||||||
+bind(page_table, total_len) KvcacheView
|
|
||||||
}
|
}
|
||||||
|
|
||||||
class KvcacheView {
|
class PageCache {
|
||||||
|
+int page_size
|
||||||
|
-PagePool _pool
|
||||||
|
-Storage _storage
|
||||||
|
-TaskTable _table
|
||||||
|
+task_alloc(task_id, prompt_ids) bool
|
||||||
|
+task_free(task_id)
|
||||||
|
+task_extend(task_id, pos) bool
|
||||||
|
+task_cached(task_id) int
|
||||||
|
+task_record_hashes(task_id, prompt_ids, start_logical_page)
|
||||||
|
+bind_tasks(task_ids, total_len, device) PageCacheView
|
||||||
|
}
|
||||||
|
|
||||||
|
class ContiguousCache {
|
||||||
|
+int max_seq_len
|
||||||
|
+Tensor k, v
|
||||||
|
+task_alloc(task_id, prompt_ids) bool
|
||||||
|
+task_free(task_id)
|
||||||
|
+task_extend(task_id, pos) bool
|
||||||
|
+bind_tasks(task_ids, total_len, device) ContiguousCacheView
|
||||||
|
}
|
||||||
|
|
||||||
|
class CacheView {
|
||||||
|
<<abstract>>
|
||||||
|
+write(layer_id, k, v)
|
||||||
|
+gather(layer_id) Tuple[Tensor, Tensor]
|
||||||
|
}
|
||||||
|
|
||||||
|
class PageCacheView {
|
||||||
-Storage _storage
|
-Storage _storage
|
||||||
+Tensor _page_table
|
+Tensor _page_table
|
||||||
+int _total_len
|
+int _total_len
|
||||||
@@ -705,6 +725,14 @@ classDiagram
|
|||||||
+gather(layer_id) Tuple[Tensor, Tensor]
|
+gather(layer_id) Tuple[Tensor, Tensor]
|
||||||
}
|
}
|
||||||
|
|
||||||
|
class ContiguousCacheView {
|
||||||
|
-ContiguousCache _cache
|
||||||
|
+Tensor _batch_indices
|
||||||
|
+int _total_len
|
||||||
|
+write(layer_id, k, v)
|
||||||
|
+gather(layer_id) Tuple[Tensor, Tensor]
|
||||||
|
}
|
||||||
|
|
||||||
class TaskTable {
|
class TaskTable {
|
||||||
+set(task_id, page_table, cached)
|
+set(task_id, page_table, cached)
|
||||||
+get(task_id) List[int]
|
+get(task_id) List[int]
|
||||||
@@ -714,23 +742,22 @@ classDiagram
|
|||||||
+table_tensor(task_ids, device) Tensor
|
+table_tensor(task_ids, device) Tensor
|
||||||
}
|
}
|
||||||
|
|
||||||
class Task {
|
class Task {
|
||||||
+str task_id
|
+str task_id
|
||||||
+List prompt_ids
|
+List prompt_ids
|
||||||
+Optional[int] max_tokens
|
+Optional[int] max_tokens
|
||||||
+float temperature
|
+float temperature
|
||||||
+float top_p
|
+float top_p
|
||||||
+int top_k
|
+int top_k
|
||||||
+TaskStatus status
|
+TaskStatus status
|
||||||
+List output_ids
|
+List output_ids
|
||||||
+int input_tokens
|
+int input_tokens
|
||||||
+int output_tokens
|
+int output_tokens
|
||||||
+float arrival_time
|
+float arrival_time
|
||||||
+Optional[float] finish_time
|
+Optional[float] finish_time
|
||||||
+Optional[Callable] stream_callback
|
+int next_pos
|
||||||
+int next_pos
|
+is_finished(stop_ids) bool
|
||||||
+is_finished(stop_ids) bool
|
}
|
||||||
}
|
|
||||||
|
|
||||||
class TaskStatus {
|
class TaskStatus {
|
||||||
<<enumeration>>
|
<<enumeration>>
|
||||||
@@ -1035,14 +1062,14 @@ classDiagram
|
|||||||
TrainCallback <|-- GradientCheckpointingCallback
|
TrainCallback <|-- GradientCheckpointingCallback
|
||||||
TrainCallback <|-- CheckpointCallback
|
TrainCallback <|-- CheckpointCallback
|
||||||
TrainCallback <|-- ProgressBarCallback
|
TrainCallback <|-- ProgressBarCallback
|
||||||
TrainCallback <|-- MetricLoggerCallback
|
TrainCallback <|-- MetricCallback
|
||||||
TrainCallback <|-- ValidationCallback
|
|
||||||
BaseDataset <|-- SEQDataset
|
BaseDataset <|-- SEQDataset
|
||||||
BaseDataset <|-- SFTDataset
|
BaseDataset <|-- SFTDataset
|
||||||
BaseDataset <|-- DPODataset
|
BaseDataset <|-- DPODataset
|
||||||
BaseDataset <|-- GRPODataset
|
BaseDataset <|-- GRPODataset
|
||||||
Store <|-- H5Store
|
Store <|-- H5Store
|
||||||
Store <|-- MmapStore
|
Store <|-- MmapStore
|
||||||
|
Store <|-- JsonlStore
|
||||||
BaseSamplingStrategy <|-- TemperatureStrategy
|
BaseSamplingStrategy <|-- TemperatureStrategy
|
||||||
BaseSamplingStrategy <|-- TopKStrategy
|
BaseSamplingStrategy <|-- TopKStrategy
|
||||||
BaseSamplingStrategy <|-- TopPStrategy
|
BaseSamplingStrategy <|-- TopPStrategy
|
||||||
@@ -1075,11 +1102,15 @@ classDiagram
|
|||||||
ResponseBuilder <|-- OpenAIResponseBuilder
|
ResponseBuilder <|-- OpenAIResponseBuilder
|
||||||
ResponseBuilder <|-- AnthropicResponseBuilder
|
ResponseBuilder <|-- AnthropicResponseBuilder
|
||||||
BaseMaskBuilder <|-- SectionedMaskBuilder
|
BaseMaskBuilder <|-- SectionedMaskBuilder
|
||||||
|
KVCache <|-- PageCache
|
||||||
|
KVCache <|-- ContiguousCache
|
||||||
|
CacheView <|-- PageCacheView
|
||||||
|
CacheView <|-- ContiguousCacheView
|
||||||
|
|
||||||
%% --- Composition (strong ownership, part destroyed with whole) ---
|
%% --- Composition (strong ownership, part destroyed with whole) ---
|
||||||
KVCache *-- PagePool
|
PageCache *-- PagePool
|
||||||
KVCache *-- Storage
|
PageCache *-- Storage
|
||||||
KVCache *-- TaskTable
|
PageCache *-- TaskTable
|
||||||
InferenceEngine *-- InferenceScheduler
|
InferenceEngine *-- InferenceScheduler
|
||||||
InferenceScheduler *-- KVCache
|
InferenceScheduler *-- KVCache
|
||||||
InferenceScheduler *-- Executor
|
InferenceScheduler *-- Executor
|
||||||
@@ -1107,7 +1138,8 @@ classDiagram
|
|||||||
TrainContext o-- BaseScheduler
|
TrainContext o-- BaseScheduler
|
||||||
TrainContext o-- Checkpoint
|
TrainContext o-- Checkpoint
|
||||||
TrainContext o-- BaseExecutor
|
TrainContext o-- BaseExecutor
|
||||||
KvcacheView o-- Storage
|
PageCacheView o-- Storage
|
||||||
|
ContiguousCacheView o-- ContiguousCache
|
||||||
SamplingPipeline o-- BaseSamplingStrategy
|
SamplingPipeline o-- BaseSamplingStrategy
|
||||||
BaseDataset o-- Store
|
BaseDataset o-- Store
|
||||||
Pipeline o-- PipelineConfig
|
Pipeline o-- PipelineConfig
|
||||||
@@ -1129,6 +1161,7 @@ classDiagram
|
|||||||
DecoderBlock ..> FFNFactory : uses
|
DecoderBlock ..> FFNFactory : uses
|
||||||
StoreFactory ..> H5Store : creates
|
StoreFactory ..> H5Store : creates
|
||||||
StoreFactory ..> MmapStore : creates
|
StoreFactory ..> MmapStore : creates
|
||||||
|
StoreFactory ..> JsonlStore : creates
|
||||||
ConfigFactory ..> AutoRegressiveLMConfig : creates
|
ConfigFactory ..> AutoRegressiveLMConfig : creates
|
||||||
ConfigFactory ..> EncoderConfig : creates
|
ConfigFactory ..> EncoderConfig : creates
|
||||||
ExecutorFactory ..> NoneExecutor : creates
|
ExecutorFactory ..> NoneExecutor : creates
|
||||||
@@ -1142,7 +1175,8 @@ classDiagram
|
|||||||
TrainContextBuilder ..> ResumableDistributedSampler : creates
|
TrainContextBuilder ..> ResumableDistributedSampler : creates
|
||||||
Checkpoint ..> Checkpoint : serializes
|
Checkpoint ..> Checkpoint : serializes
|
||||||
CheckpointCallback ..> Checkpoint : creates
|
CheckpointCallback ..> Checkpoint : creates
|
||||||
KVCache ..> KvcacheView : binds
|
PageCache ..> PageCacheView : binds
|
||||||
|
ContiguousCache ..> ContiguousCacheView : binds
|
||||||
InferenceEngine ..> GenerationRequest : uses
|
InferenceEngine ..> GenerationRequest : uses
|
||||||
InferenceEngine ..> GenerateResult : creates
|
InferenceEngine ..> GenerateResult : creates
|
||||||
OpenAIResponseBuilder ..> ChatCompletionRequest : receives
|
OpenAIResponseBuilder ..> ChatCompletionRequest : receives
|
||||||
@@ -1171,12 +1205,12 @@ classDiagram
|
|||||||
|--------|------------|-------------|
|
|--------|------------|-------------|
|
||||||
| **astrai.config** | BaseConfig, BaseModelConfig, AutoRegressiveLMConfig, EncoderConfig, ConfigFactory, TrainConfig, PipelineConfig, InputConfig, ProcessingConfig, OutputConfig | Configuration management (to_dict/from_dict, to_file/from_file) |
|
| **astrai.config** | BaseConfig, BaseModelConfig, AutoRegressiveLMConfig, EncoderConfig, ConfigFactory, TrainConfig, PipelineConfig, InputConfig, ProcessingConfig, OutputConfig | Configuration management (to_dict/from_dict, to_file/from_file) |
|
||||||
| **astrai.preprocessing** | BaseMaskBuilder, MaskBuilderFactory, SectionedMaskBuilder, Pipeline, filter_by_length, PackingStrategy, PackingStrategyFactory, PositionIdStrategy, PositionIdStrategyFactory, StoreWriter, StoreWriterFactory | Declarative JSON-driven data preprocessing |
|
| **astrai.preprocessing** | BaseMaskBuilder, MaskBuilderFactory, SectionedMaskBuilder, Pipeline, filter_by_length, PackingStrategy, PackingStrategyFactory, PositionIdStrategy, PositionIdStrategyFactory, StoreWriter, StoreWriterFactory | Declarative JSON-driven data preprocessing |
|
||||||
| **astrai.dataset** | BaseDataset–GRPODataset, Store–MmapStore, StoreFactory, ResumableDistributedSampler, DatasetFactory | Dataset loading and management |
|
| **astrai.dataset** | BaseDataset–GRPODataset, Store–JsonlStore/MmapStore/H5Store, StoreFactory, ResumableDistributedSampler, 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, 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, BaseStrategy–GRPOStrategy, StrategyFactory, BaseScheduler–WSDScheduler, SchedulerFactory, TrainCallback(Protocol)–ValidationCallback, CallbackFactory | Training workflow |
|
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy–GRPOStrategy, StrategyFactory, BaseScheduler–WSDScheduler, SchedulerFactory, TrainCallback(Protocol)–MetricCallback, CallbackFactory | Training workflow |
|
||||||
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, KVCache–KvcacheView, Allocator–Storage, Task, TaskManager, TaskStatus, GenerationRequest, GenerateResult, BaseSamplingStrategy–SamplingPipeline, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, ChatMessage–MessagesRequest, app | Inference service |
|
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, KVCache–ContiguousCache/PageCache, CacheView–ContiguousCacheView/PageCacheView, Allocator–Storage, Task, TaskManager, TaskStatus, GenerationRequest, GenerateResult, BaseSamplingStrategy–SamplingPipeline, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, ChatMessage–MessagesRequest, app | 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, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, FSDPExecutor, 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 |
|
||||||
@@ -1195,7 +1229,7 @@ classDiagram
|
|||||||
| **Context** | `TrainContext` | Unified training state bag |
|
| **Context** | `TrainContext` | Unified training state bag |
|
||||||
| **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction |
|
| **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction |
|
||||||
| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor` | Gradient accumulation & model distribution |
|
| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor` | Gradient accumulation & model distribution |
|
||||||
| **Storage** | `Store`, `H5Store`, `MmapStore` | 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 |
|
||||||
|
|
||||||
@@ -1207,10 +1241,10 @@ classDiagram
|
|||||||
4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)` → `NoneExecutor` / `DDPExecutor` / `FSDPExecutor`
|
4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)` → `NoneExecutor` / `DDPExecutor` / `FSDPExecutor`
|
||||||
5. **Inference Flow**: `InferenceEngine` → `InferenceScheduler` → `AutoRegressiveLM`, backed by `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) 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`
|
||||||
8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata (rank-0 only), extra state saved as `{key}.pt`
|
8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata (rank-0 only), extra state saved as `{key}.pt`
|
||||||
9. **Scheduler**: `SchedulerFactory` creates `CosineScheduler`/`SGDRScheduler`/`WSDScheduler`
|
9. **Scheduler**: `SchedulerFactory` creates `CosineScheduler`/`SGDRScheduler`/`WSDScheduler`
|
||||||
10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops
|
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-05-30
|
> Document Update Time: 2026-07-09
|
||||||
|
|||||||
+10
-7
@@ -48,23 +48,26 @@ The output `meta.json` records the storage format, key names, dtype, total token
|
|||||||
|
|
||||||
`detect_format(load_path)` inspects the path:
|
`detect_format(load_path)` inspects the path:
|
||||||
|
|
||||||
- If `load_path` is a file: checks suffix — `.h5`/`.hdf5` → `"h5"`, unknown suffix raises `ValueError`
|
- If `load_path` is a file: checks suffix — `.h5`/`.hdf5` → `"h5"`, `.jsonl` → `"jsonl"`, unknown suffix raises `ValueError`
|
||||||
- If `load_path` is a directory: recursively globs for `*.h5`/`*.hdf5` files → `"h5"`, or `*.bin` + `**/meta.json` → `"bin"`
|
- If `load_path` is a directory: recursively globs for `*.h5`/`*.hdf5` files → `"h5"`, `*.bin` + `**/meta.json` → `"bin"`, or `*.jsonl` + `dataset_config.json` → `"jsonl"`
|
||||||
|
|
||||||
### Store Backends
|
### Store Backends
|
||||||
|
|
||||||
Storage format is auto-detected by `detect_format()`; backends are dispatched via registry:
|
Storage format is auto-detected by `detect_format()`; backends are dispatched via registry:
|
||||||
|
|
||||||
```
|
```
|
||||||
StoreFactory.create("h5") → H5Store
|
StoreFactory.create("h5") → H5Store
|
||||||
StoreFactory.create("bin") → MmapStore
|
StoreFactory.create("bin") → MmapStore
|
||||||
|
StoreFactory.create("jsonl") → JsonlStore
|
||||||
```
|
```
|
||||||
|
|
||||||
**H5Store**: Reads HDF5 files, supports `share_memory_()` for multi-process DataLoader workers (copies tensors to shared memory).
|
**H5Store**: Reads HDF5 files. Tensors are loaded into host memory and normalized into segmented storage.
|
||||||
|
|
||||||
**MmapStore**: Memory-maps `.bin` files. OS page cache sharing is native — no explicit `share_memory_()` needed. Uses `torch.from_numpy(np.memmap(...))`.
|
**MmapStore**: Memory-maps `.bin` files. OS page cache sharing is native — no explicit `share_memory_()` needed. Uses `torch.from_numpy(np.memmap(...))`.
|
||||||
|
|
||||||
Both backends normalise tensors into `Store._data[Dict[str, List[Tensor]]]` + `Store._cum[Dict[str, List[int]]]` (cumulative lengths for bisect-based indexing).
|
**JsonlStore**: On-the-fly tokenization of raw JSONL files at load time. Requires a `dataset_config.json` alongside the `.jsonl` files following the same `PipelineConfig` schema with an additional `tokenizer_path` field.
|
||||||
|
|
||||||
|
All backends normalise tensors into `Store._data[Dict[str, List[Tensor]]]` + `Store._cum[Dict[str, List[int]]]` (cumulative lengths for bisect-based indexing).
|
||||||
|
|
||||||
## Data Keys by Training Type
|
## Data Keys by Training Type
|
||||||
|
|
||||||
@@ -106,4 +109,4 @@ DatasetFactory.load(train_type, load_path, window_size, stride=None, storage_typ
|
|||||||
|
|
||||||
Standard PyTorch `DataLoader` with configurable `batch_size`, `num_workers`, `pin_memory`, `prefetch_factor`. Sampler produces indices; dataloader fetches tensor batches via `__getitem__`.
|
Standard PyTorch `DataLoader` with configurable `batch_size`, `num_workers`, `pin_memory`, `prefetch_factor`. Sampler produces indices; dataloader fetches tensor batches via `__getitem__`.
|
||||||
|
|
||||||
> Document Update Time: 2026-06-19
|
> Document Update Time: 2026-07-09
|
||||||
|
|||||||
+24
-13
@@ -23,29 +23,40 @@ RoPE is applied **before** KV cache write, not after — otherwise position enco
|
|||||||
|
|
||||||
## KVCache System
|
## KVCache System
|
||||||
|
|
||||||
Seven classes working together:
|
Seven classes working together, with two concrete cache implementations:
|
||||||
|
|
||||||
|
### ContiguousCache (default)
|
||||||
|
|
||||||
```
|
```
|
||||||
KVCache (facade)
|
ContiguousCache (simple contiguous per-slot cache)
|
||||||
├── PagePool orchestrates page allocation + prefix matching
|
├── ContiguousCacheView bundles k/v tensors + slot indices for attention layers
|
||||||
│ ├── Allocator bitmask-based page allocator + ref-count + LRU eviction (inside PagePool)
|
|
||||||
│ └── PrefixCache hash-based prefix matching (page_hash via polynomial hash) (inside PagePool)
|
|
||||||
├── 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)
|
|
||||||
└── KvcacheView bundles Storage + page_table + total_len for attention layers (returned by bind())
|
|
||||||
```
|
```
|
||||||
|
|
||||||
`KVCache.bind(page_table, total_len)` returns a `KvcacheView` used by attention layers via `write()` / `gather()`.
|
Created by default when no cache is passed to `InferenceScheduler`. Each task occupies a fixed slot of `[max_seq_len, n_kv_heads, head_dim]`. Simple and efficient for small-to-medium batch sizes.
|
||||||
|
|
||||||
|
### PageCache (paged with prefix sharing)
|
||||||
|
|
||||||
|
```
|
||||||
|
PageCache (paged KV cache with prefix sharing, alternative)
|
||||||
|
├── PagePool orchestrates page allocation + prefix matching
|
||||||
|
│ ├── Allocator bitmask-based page allocator + ref-count + LRU
|
||||||
|
│ └── PrefixCache hash-based prefix matching (page_hash via polynomial hash)
|
||||||
|
├── TaskTable maps task_id → page_table + cached token count
|
||||||
|
├── Storage k_cache / v_cache tensors (n_layers × n_pages × page_size × n_kv_heads × head_dim)
|
||||||
|
└── PageCacheView bundles Storage + page_table + total_len for attention layers
|
||||||
|
```
|
||||||
|
|
||||||
|
`isinstance(cache, KVCache)` checks dispatch to the correct view. Both implement the abstract `KVCache` interface used by `Executor` and `InferenceScheduler`.
|
||||||
|
|
||||||
## Continuous Batching
|
## Continuous Batching
|
||||||
|
|
||||||
`InferenceScheduler` runs a daemon thread with a 4-phase loop:
|
`InferenceScheduler` runs a daemon thread with a 4-phase loop:
|
||||||
|
|
||||||
```
|
```
|
||||||
1. Cleanup → Remove finished tasks, free KV pages
|
1. Cleanup → Remove finished tasks, free KV cache slots/pages
|
||||||
2. Refill → Pop from waiting_queue, task_alloc pages, activate
|
2. Refill → Pop from waiting_queue, task_alloc resources, activate
|
||||||
3. Prefill → Group by (prompt_len, start_pos), run full forward
|
3. Prefill → Group by (prompt_len, start_pos), run full forward
|
||||||
4. Decode → Pick largest same-position group, single-token forward
|
4. Decode → Run single-token forward for each same-position group
|
||||||
```
|
```
|
||||||
|
|
||||||
## Sampling (Strategy Pattern)
|
## Sampling (Strategy Pattern)
|
||||||
@@ -238,4 +249,4 @@ async for token in engine.generate_async("Hello", ...): # -> AsyncGenerator[s
|
|||||||
print(token)
|
print(token)
|
||||||
```
|
```
|
||||||
|
|
||||||
> Document Update Time: 2026-06-19
|
> Document Update Time: 2026-07-09
|
||||||
|
|||||||
+11
-10
@@ -28,13 +28,17 @@
|
|||||||
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
|
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
|
||||||
| `--max_grad_norm` | Maximum gradient norm for clipping | 1.0 |
|
| `--max_grad_norm` | Maximum gradient norm for clipping | 1.0 |
|
||||||
|
|
||||||
### Optimizer (AdamW)
|
### Optimizer (MuonMix)
|
||||||
|
|
||||||
|
Combined optimizer: matrix parameters via **Muon**, non-matrix via **AdamW** (`fused=True`).
|
||||||
|
|
||||||
| Parameter | Description | Default |
|
| Parameter | Description | Default |
|
||||||
|-----------|-------------|---------|
|
|-----------|-------------|---------|
|
||||||
| `--adamw_beta1` | AdamW beta1 | 0.9 |
|
| `--weight_decay` | Weight decay (applied to Muon matrix params; non-matrix use 0) | 0.1 |
|
||||||
| `--adamw_beta2` | AdamW beta2 | 0.95 |
|
| `--muon_momentum` | Muon momentum factor | 0.95 |
|
||||||
| `--adamw_weight_decay` | AdamW weight decay | 0.01 |
|
| `--muon_nesterov` | Enable Nesterov momentum for Muon | True |
|
||||||
|
| `--muon_ns_steps` | Newton-Schulz iteration steps for Muon | 5 |
|
||||||
|
| `--muon_adjust_lr` | Muon LR adjustment strategy (`original`, `match_rms_adamw`) | `match_rms_adamw` |
|
||||||
|
|
||||||
### Data Loading
|
### Data Loading
|
||||||
|
|
||||||
@@ -67,7 +71,6 @@
|
|||||||
| Parameter | Description | Default |
|
| Parameter | Description | Default |
|
||||||
|-----------|-------------|---------|
|
|-----------|-------------|---------|
|
||||||
| `--log_dir` | Directory for metric logs | checkpoint/logs |
|
| `--log_dir` | Directory for metric logs | checkpoint/logs |
|
||||||
| `--log_interval` | Number of optimizer steps between metric logs | 1 |
|
|
||||||
| `--metrics` | Metrics to log (e.g. --metrics loss lr val_loss) | ["loss", "lr", "grad_norm"] |
|
| `--metrics` | Metrics to log (e.g. --metrics loss lr val_loss) | ["loss", "lr", "grad_norm"] |
|
||||||
|
|
||||||
### Gradient Checkpointing
|
### Gradient Checkpointing
|
||||||
@@ -105,7 +108,7 @@
|
|||||||
| 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) |
|
| `--min_rate` | Minimum LR as fraction of base LR | None (scheduler default: 0.01) |
|
||||||
| `--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 (required for wsd) |
|
||||||
@@ -127,9 +130,7 @@ nohup python scripts/tools/train.py \
|
|||||||
--warmup_ratio=0.05 \
|
--warmup_ratio=0.05 \
|
||||||
--max_lr=1e-4 \
|
--max_lr=1e-4 \
|
||||||
--max_grad_norm=1.0 \
|
--max_grad_norm=1.0 \
|
||||||
--adamw_beta1=0.9 \
|
--weight_decay=0.1 \
|
||||||
--adamw_beta2=0.95 \
|
|
||||||
--adamw_weight_decay=0.01 \
|
|
||||||
--window_size=2048 \
|
--window_size=2048 \
|
||||||
--ckpt_interval=10000 \
|
--ckpt_interval=10000 \
|
||||||
--ckpt_dir=./checkpoint \
|
--ckpt_dir=./checkpoint \
|
||||||
@@ -200,4 +201,4 @@ See [Preprocessing Guide](preprocessing.md) for config file format and examples.
|
|||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
> Document Update Time: 2026-06-19
|
> Document Update Time: 2026-07-09
|
||||||
@@ -268,7 +268,7 @@ When `sources` is set, `sections` is ignored.
|
|||||||
| `storage_format` | str | `"bin"` | `"bin"` (mmap) or `"h5"` |
|
| `storage_format` | str | `"bin"` | `"bin"` (mmap) or `"h5"` |
|
||||||
| `max_tokens_per_shard` | int | `100000000` | Flush threshold in cumulative tokens |
|
| `max_tokens_per_shard` | int | `100000000` | Flush threshold in cumulative tokens |
|
||||||
| `dtype` | dict[str, str] | `{}` | Per-key tensor dtype override (e.g. `{"loss_mask": "bool"}`) |
|
| `dtype` | dict[str, str] | `{}` | Per-key tensor dtype override (e.g. `{"loss_mask": "bool"}`) |
|
||||||
| `position_ids_mode` | str | `"none"` | How to compute position_ids: `"none"`, `"doc_reset"`, `"continuous"` |
|
| `position_ids_mode` | str | `"doc_reset"` | How to compute position_ids: `"none"`, `"doc_reset"`, `"continuous"` |
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
@@ -361,4 +361,4 @@ Pipeline(
|
|||||||
).run()
|
).run()
|
||||||
```
|
```
|
||||||
|
|
||||||
> Document Update Time: 2026-06-03
|
> Document Update Time: 2026-07-09
|
||||||
|
|||||||
+9
-11
@@ -80,13 +80,13 @@ on_train_end
|
|||||||
| `on_train_begin` | Before training starts | `GradientCheckpointingCallback` |
|
| `on_train_begin` | Before training starts | `GradientCheckpointingCallback` |
|
||||||
| `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` |
|
| `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` |
|
||||||
| `on_batch_begin` | Every batch | — |
|
| `on_batch_begin` | Every batch | — |
|
||||||
| `on_optimizer_step` | Every accumulation window | `GradientClippingCallback`, `MetricLoggerCallback`, `ValidationCallback` |
|
| `on_optimizer_step` | Every accumulation window | `GradientClippingCallback`, `MetricCallback`, `ProgressBarCallback` |
|
||||||
| `on_batch_end` | Every batch | `CheckpointCallback`, `MetricLoggerCallback`, `ProgressBarCallback` |
|
| `on_batch_end` | Every batch | `CheckpointCallback` |
|
||||||
| `on_epoch_end` | End of each epoch | `ProgressBarCallback` |
|
| `on_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` |
|
||||||
| `on_error` | On exception during training | `CheckpointCallback`, `MetricLoggerCallback` |
|
| `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` |
|
||||||
| `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricLoggerCallback`, `GradientCheckpointingCallback` |
|
| `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricCallback`, `GradientCheckpointingCallback` |
|
||||||
|
|
||||||
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `validation` (periodic validation on val_dataset), `metric_logger` (JSONL, rank-0), `progress_bar` (tqdm), `gradient_clipping`.
|
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm), `gradient_clipping`.
|
||||||
|
|
||||||
## Strategies
|
## Strategies
|
||||||
|
|
||||||
@@ -108,7 +108,7 @@ $$
|
|||||||
L_{\text{SFT}} = -\sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta)
|
L_{\text{SFT}} = -\sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta)
|
||||||
$$
|
$$
|
||||||
|
|
||||||
Keys: `input_ids`, `target_ids`, `loss_mask`. Optional: `label_smoothing`.
|
Keys: `input_ids`, `target_ids`, `loss_mask`, `position_ids`. Optional: `label_smoothing`.
|
||||||
|
|
||||||
### DPO (Direct Preference Optimization)
|
### DPO (Direct Preference Optimization)
|
||||||
|
|
||||||
@@ -201,9 +201,7 @@ nohup python scripts/tools/train.py \
|
|||||||
--warmup_ratio=0.05 \
|
--warmup_ratio=0.05 \
|
||||||
--max_lr=1e-4 \
|
--max_lr=1e-4 \
|
||||||
--max_grad_norm=1.0 \
|
--max_grad_norm=1.0 \
|
||||||
--adamw_beta1=0.9 \
|
--weight_decay=0.1 \
|
||||||
--adamw_beta2=0.95 \
|
|
||||||
--adamw_weight_decay=0.01 \
|
|
||||||
--window_size=2048 \
|
--window_size=2048 \
|
||||||
--ckpt_interval=10000 \
|
--ckpt_interval=10000 \
|
||||||
--ckpt_dir=./checkpoint \
|
--ckpt_dir=./checkpoint \
|
||||||
@@ -214,4 +212,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-05-30
|
> Document Update Time: 2026-07-09
|
||||||
|
|||||||
+1
-1
@@ -1,4 +1,4 @@
|
|||||||
__version__ = "1.3.7"
|
__version__ = "1.3.9"
|
||||||
__author__ = "ViperEkura"
|
__author__ = "ViperEkura"
|
||||||
|
|
||||||
from astrai.config import (
|
from astrai.config import (
|
||||||
|
|||||||
@@ -96,7 +96,7 @@ class OutputConfig(BaseConfig):
|
|||||||
storage_format: str = "bin"
|
storage_format: str = "bin"
|
||||||
max_tokens_per_shard: int = 100_000_000
|
max_tokens_per_shard: int = 100_000_000
|
||||||
dtype: Dict[str, str] = field(default_factory=dict)
|
dtype: Dict[str, str] = field(default_factory=dict)
|
||||||
position_ids_mode: str = "none"
|
position_ids_mode: str = "doc_reset"
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|||||||
@@ -0,0 +1,21 @@
|
|||||||
|
"""CUDA attention kernel wrappers with torch fallback.
|
||||||
|
|
||||||
|
Public API:
|
||||||
|
- ``attn_decode`` — single-query decode attention
|
||||||
|
- ``attn_prefill`` — multi-query prefill attention
|
||||||
|
- ``attn_paged_decode`` — paged decode attention (direct page-table access)
|
||||||
|
|
||||||
|
Each wrapper dispatches to its compiled CUDA kernel (``astrai.extension.attn_*``)
|
||||||
|
when available, otherwise falls back to ``torch.nn.functional.scaled_dot_product_attention``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from astrai.extension.loader import KERNEL_NAMES, is_available
|
||||||
|
from astrai.extension.ops import attn_decode, attn_paged_decode, attn_prefill
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"attn_decode",
|
||||||
|
"attn_paged_decode",
|
||||||
|
"attn_prefill",
|
||||||
|
"is_available",
|
||||||
|
"KERNEL_NAMES",
|
||||||
|
]
|
||||||
@@ -0,0 +1,36 @@
|
|||||||
|
"""Dynamic discovery and loading of compiled CUDA kernel modules.
|
||||||
|
|
||||||
|
Each kernel is registered in ``csrc/build.py`` and built into a ``.so`` placed
|
||||||
|
in this package directory. On import we try to load each one; kernels that
|
||||||
|
failed to build (or are running on a CPU-only machine) are marked unavailable
|
||||||
|
so the wrapper functions can fall back to ``torch`` SDPA.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import importlib
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
KERNEL_NAMES = ["attn_decode", "attn_prefill", "attn_paged_decode"]
|
||||||
|
|
||||||
|
_available: dict[str, bool] = {}
|
||||||
|
_modules: dict[str, object] = {}
|
||||||
|
|
||||||
|
for _name in KERNEL_NAMES:
|
||||||
|
try:
|
||||||
|
_mod = importlib.import_module(f".{_name}", package=__package__)
|
||||||
|
_available[_name] = True
|
||||||
|
_modules[_name] = _mod
|
||||||
|
except ImportError:
|
||||||
|
_available[_name] = False
|
||||||
|
_modules[_name] = None
|
||||||
|
|
||||||
|
|
||||||
|
def is_available(name: str) -> bool:
|
||||||
|
"""Return ``True`` if the compiled kernel ``name`` was loaded."""
|
||||||
|
return _available.get(name, False)
|
||||||
|
|
||||||
|
|
||||||
|
def get_module(name: str) -> object:
|
||||||
|
"""Return the loaded kernel module for ``name``, or ``None`` if unavailable."""
|
||||||
|
return _modules.get(name)
|
||||||
@@ -0,0 +1,149 @@
|
|||||||
|
"""GQA attention wrapper functions — one entry point per compiled kernel.
|
||||||
|
|
||||||
|
Each wrapper dispatches to its CUDA kernel (loaded in ``loader.py``) when
|
||||||
|
available, otherwise falls back to ``torch`` SDPA.
|
||||||
|
|
||||||
|
Add new kernel wrappers here; split into per-variant files only if this file
|
||||||
|
grows large.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
from astrai.extension.loader import _available, _modules
|
||||||
|
|
||||||
|
|
||||||
|
def _expand_kv_heads(
|
||||||
|
k: torch.Tensor, v: torch.Tensor, q_head: int
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""Expand K/V heads to match Q heads for GQA fallback."""
|
||||||
|
kv_head = k.size(1)
|
||||||
|
if kv_head == q_head:
|
||||||
|
return k, v
|
||||||
|
group = q_head // kv_head
|
||||||
|
k = k.repeat_interleave(group, dim=1)
|
||||||
|
v = v.repeat_interleave(group, dim=1)
|
||||||
|
return k, v
|
||||||
|
|
||||||
|
|
||||||
|
def _torch_fallback(
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
v: torch.Tensor,
|
||||||
|
mask: torch.Tensor | None,
|
||||||
|
is_causal: bool,
|
||||||
|
scale: float | None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Reference attention via ``scaled_dot_product_attention``."""
|
||||||
|
k, v = _expand_kv_heads(k, v, q.size(1))
|
||||||
|
attn_mask = mask[:, None, None, :] if mask is not None else None
|
||||||
|
return F.scaled_dot_product_attention(
|
||||||
|
q, k, v, attn_mask=attn_mask, is_causal=is_causal and mask is None, scale=scale
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _gather_kv_from_pages(
|
||||||
|
page_table: torch.Tensor,
|
||||||
|
k_cache: torch.Tensor,
|
||||||
|
v_cache: torch.Tensor,
|
||||||
|
page_size: int,
|
||||||
|
kv_len: int,
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""Gather contiguous K/V from paged cache for torch SDPA fallback.
|
||||||
|
|
||||||
|
Shapes:
|
||||||
|
page_table : [batch, max_pages] (int64)
|
||||||
|
k_cache : [n_pages, page_size, n_kv_heads, head_dim]
|
||||||
|
v_cache : same as k_cache
|
||||||
|
Returns:
|
||||||
|
k, v : [batch, n_kv_heads, kv_len, head_dim]
|
||||||
|
"""
|
||||||
|
batch, max_pages = page_table.shape
|
||||||
|
n_pages, ps, n_kv_heads, head_dim = k_cache.shape
|
||||||
|
if ps != page_size:
|
||||||
|
raise ValueError(f"k_cache page_size mismatch: {ps} vs {page_size}")
|
||||||
|
|
||||||
|
k = k_cache.new_empty(batch, n_kv_heads, kv_len, head_dim)
|
||||||
|
v = v_cache.new_empty(batch, n_kv_heads, kv_len, head_dim)
|
||||||
|
|
||||||
|
for b in range(batch):
|
||||||
|
for pos in range(kv_len):
|
||||||
|
log_pg = pos // page_size
|
||||||
|
pg_off = pos % page_size
|
||||||
|
phys = int(page_table[b, log_pg].item())
|
||||||
|
k[b, :, pos, :] = k_cache[phys, pg_off, :, :]
|
||||||
|
v[b, :, pos, :] = v_cache[phys, pg_off, :, :]
|
||||||
|
return k, v
|
||||||
|
|
||||||
|
|
||||||
|
def attn_decode(
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
v: torch.Tensor,
|
||||||
|
mask: torch.Tensor | None = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
causal_offset: int = 0,
|
||||||
|
scale: float | None = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if _available["attn_decode"]:
|
||||||
|
return _modules["attn_decode"].attn_decode(
|
||||||
|
q,
|
||||||
|
k,
|
||||||
|
v,
|
||||||
|
mask=mask,
|
||||||
|
is_causal=is_causal,
|
||||||
|
causal_offset=causal_offset,
|
||||||
|
scale=scale,
|
||||||
|
)
|
||||||
|
return _torch_fallback(q, k, v, mask, is_causal, scale)
|
||||||
|
|
||||||
|
|
||||||
|
def attn_prefill(
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
v: torch.Tensor,
|
||||||
|
mask: torch.Tensor | None = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
causal_offset: int = 0,
|
||||||
|
scale: float | None = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if _available["attn_prefill"]:
|
||||||
|
return _modules["attn_prefill"].attn_prefill(
|
||||||
|
q,
|
||||||
|
k,
|
||||||
|
v,
|
||||||
|
mask=mask,
|
||||||
|
is_causal=is_causal,
|
||||||
|
causal_offset=causal_offset,
|
||||||
|
scale=scale,
|
||||||
|
)
|
||||||
|
return _torch_fallback(q, k, v, mask, is_causal, scale)
|
||||||
|
|
||||||
|
|
||||||
|
def attn_paged_decode(
|
||||||
|
q: torch.Tensor,
|
||||||
|
page_table: torch.Tensor,
|
||||||
|
k_cache: torch.Tensor,
|
||||||
|
v_cache: torch.Tensor,
|
||||||
|
page_size: int,
|
||||||
|
kv_len: int,
|
||||||
|
mask: torch.Tensor | None = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
causal_offset: int = 0,
|
||||||
|
scale: float | None = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if _available["attn_paged_decode"]:
|
||||||
|
return _modules["attn_paged_decode"].attn_paged_decode(
|
||||||
|
q,
|
||||||
|
page_table,
|
||||||
|
k_cache,
|
||||||
|
v_cache,
|
||||||
|
page_size,
|
||||||
|
kv_len,
|
||||||
|
mask=mask,
|
||||||
|
is_causal=is_causal,
|
||||||
|
causal_offset=causal_offset,
|
||||||
|
scale=scale,
|
||||||
|
)
|
||||||
|
k, v = _gather_kv_from_pages(page_table, k_cache, v_cache, page_size, kv_len)
|
||||||
|
return _torch_fallback(q, k, v, mask, is_causal, scale)
|
||||||
@@ -30,10 +30,14 @@ from astrai.inference.api.openai import OpenAIResponseBuilder
|
|||||||
from astrai.inference.core import (
|
from astrai.inference.core import (
|
||||||
STOP,
|
STOP,
|
||||||
Allocator,
|
Allocator,
|
||||||
|
CacheView,
|
||||||
|
ContiguousCache,
|
||||||
|
ContiguousCacheView,
|
||||||
Executor,
|
Executor,
|
||||||
InferenceScheduler,
|
InferenceScheduler,
|
||||||
KVCache,
|
KVCache,
|
||||||
KvcacheView,
|
PageCache,
|
||||||
|
PageCacheView,
|
||||||
PagePool,
|
PagePool,
|
||||||
PrefixCache,
|
PrefixCache,
|
||||||
Storage,
|
Storage,
|
||||||
@@ -63,8 +67,12 @@ __all__ = [
|
|||||||
"TaskManager",
|
"TaskManager",
|
||||||
"TaskStatus",
|
"TaskStatus",
|
||||||
"Allocator",
|
"Allocator",
|
||||||
|
"CacheView",
|
||||||
"KVCache",
|
"KVCache",
|
||||||
"KvcacheView",
|
"ContiguousCache",
|
||||||
|
"ContiguousCacheView",
|
||||||
|
"PageCache",
|
||||||
|
"PageCacheView",
|
||||||
"PagePool",
|
"PagePool",
|
||||||
"PrefixCache",
|
"PrefixCache",
|
||||||
"Storage",
|
"Storage",
|
||||||
|
|||||||
@@ -2,8 +2,12 @@
|
|||||||
|
|
||||||
from astrai.inference.core.cache import (
|
from astrai.inference.core.cache import (
|
||||||
Allocator,
|
Allocator,
|
||||||
|
CacheView,
|
||||||
|
ContiguousCache,
|
||||||
|
ContiguousCacheView,
|
||||||
KVCache,
|
KVCache,
|
||||||
KvcacheView,
|
PageCache,
|
||||||
|
PageCacheView,
|
||||||
PagePool,
|
PagePool,
|
||||||
PrefixCache,
|
PrefixCache,
|
||||||
Storage,
|
Storage,
|
||||||
@@ -16,8 +20,12 @@ from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
|
|||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"Allocator",
|
"Allocator",
|
||||||
|
"CacheView",
|
||||||
"KVCache",
|
"KVCache",
|
||||||
"KvcacheView",
|
"ContiguousCache",
|
||||||
|
"ContiguousCacheView",
|
||||||
|
"PageCache",
|
||||||
|
"PageCacheView",
|
||||||
"PagePool",
|
"PagePool",
|
||||||
"PrefixCache",
|
"PrefixCache",
|
||||||
"Storage",
|
"Storage",
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
import threading
|
import threading
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
from typing import Callable, Dict, List, Optional, Tuple
|
from typing import Callable, Dict, List, Optional, Tuple
|
||||||
|
|
||||||
@@ -62,7 +63,8 @@ class Allocator:
|
|||||||
|
|
||||||
def touch(self, idx: int):
|
def touch(self, idx: int):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._lru.move_to_end(idx)
|
if idx in self._lru:
|
||||||
|
self._lru.move_to_end(idx)
|
||||||
|
|
||||||
|
|
||||||
class PrefixCache:
|
class PrefixCache:
|
||||||
@@ -274,7 +276,42 @@ class Storage:
|
|||||||
return k, v
|
return k, v
|
||||||
|
|
||||||
|
|
||||||
class KvcacheView:
|
class CacheView(ABC):
|
||||||
|
"""Abstract view passed to attention layers for KV-cache I/O."""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def write(self, layer_id: int, k: Tensor, v: Tensor): ...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]: ...
|
||||||
|
|
||||||
|
|
||||||
|
class KVCache(ABC):
|
||||||
|
"""Abstract KV-cache facade for scheduler/executor."""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool: ...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def task_free(self, task_id: str): ...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def task_extend(self, task_id: str, pos: int) -> bool: ...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def bind_tasks(
|
||||||
|
self, task_ids: List[str], total_len: int, device: torch.device
|
||||||
|
) -> CacheView: ...
|
||||||
|
|
||||||
|
def task_cached(self, task_id: str) -> int:
|
||||||
|
return 0
|
||||||
|
|
||||||
|
def task_record_hashes(
|
||||||
|
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
||||||
|
): ...
|
||||||
|
|
||||||
|
|
||||||
|
class PageCacheView(CacheView):
|
||||||
"""Bundles Storage + page_table + total_len for attention layers."""
|
"""Bundles Storage + page_table + total_len for attention layers."""
|
||||||
|
|
||||||
def __init__(self, storage: Storage, page_table: Tensor, total_len: int = 0):
|
def __init__(self, storage: Storage, page_table: Tensor, total_len: int = 0):
|
||||||
@@ -290,8 +327,8 @@ class KvcacheView:
|
|||||||
return self._storage.gather(layer_id, self._page_table, self._total_len)
|
return self._storage.gather(layer_id, self._page_table, self._total_len)
|
||||||
|
|
||||||
|
|
||||||
class KVCache:
|
class PageCache(KVCache):
|
||||||
"""Facade: page management + KV-cache I/O for continuous batching."""
|
"""Paged KV-cache with prefix sharing."""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -361,8 +398,102 @@ class KVCache:
|
|||||||
for i in range(start_logical_page, full_pages):
|
for i in range(start_logical_page, full_pages):
|
||||||
self._pool.record(page_table[i], prompt_ids, i)
|
self._pool.record(page_table[i], prompt_ids, i)
|
||||||
|
|
||||||
def make_table_tensor(self, task_ids: List[str], device: torch.device) -> Tensor:
|
def bind_tasks(
|
||||||
return self._table.table_tensor(task_ids, device)
|
self, task_ids: List[str], total_len: int, device: torch.device
|
||||||
|
) -> PageCacheView:
|
||||||
|
page_table = self._table.table_tensor(task_ids, device)
|
||||||
|
return PageCacheView(self._storage, page_table, total_len)
|
||||||
|
|
||||||
def bind(self, page_table: Tensor, total_len: int = 0) -> KvcacheView:
|
|
||||||
return KvcacheView(self._storage, page_table, total_len)
|
class ContiguousCacheView(CacheView):
|
||||||
|
"""Contiguous KV-cache view for attention layers."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self, cache: "ContiguousCache", batch_indices: Tensor, total_len: int = 0
|
||||||
|
):
|
||||||
|
self._cache = cache
|
||||||
|
self._batch_indices = batch_indices
|
||||||
|
self._total_len = total_len
|
||||||
|
|
||||||
|
def write(self, layer_id: int, k: Tensor, v: Tensor):
|
||||||
|
seq_len = k.size(1)
|
||||||
|
start_pos = self._total_len - seq_len
|
||||||
|
indices = self._batch_indices
|
||||||
|
self._cache.k[layer_id, indices, start_pos : start_pos + seq_len] = k
|
||||||
|
self._cache.v[layer_id, indices, start_pos : start_pos + seq_len] = v
|
||||||
|
new_len = start_pos + seq_len
|
||||||
|
for s in indices.tolist():
|
||||||
|
cur = self._cache._slot_len.get(s, 0)
|
||||||
|
if new_len > cur:
|
||||||
|
self._cache._slot_len[s] = new_len
|
||||||
|
|
||||||
|
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
|
||||||
|
max_len = max(
|
||||||
|
self._cache._slot_len.get(int(s), 0) for s in self._batch_indices.tolist()
|
||||||
|
)
|
||||||
|
indices = self._batch_indices
|
||||||
|
k = self._cache.k[layer_id, indices, :max_len]
|
||||||
|
v = self._cache.v[layer_id, indices, :max_len]
|
||||||
|
return k, v
|
||||||
|
|
||||||
|
|
||||||
|
class ContiguousCache(KVCache):
|
||||||
|
"""Contiguous per-slot KV cache (default implementation)."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
n_layers: int,
|
||||||
|
max_batch_size: int,
|
||||||
|
max_seq_len: int,
|
||||||
|
n_kv_heads: int,
|
||||||
|
head_dim: int,
|
||||||
|
device: torch.device,
|
||||||
|
dtype: torch.dtype,
|
||||||
|
):
|
||||||
|
self.max_seq_len = max_seq_len
|
||||||
|
self.k = torch.zeros(
|
||||||
|
n_layers,
|
||||||
|
max_batch_size,
|
||||||
|
max_seq_len,
|
||||||
|
n_kv_heads,
|
||||||
|
head_dim,
|
||||||
|
device=device,
|
||||||
|
dtype=dtype,
|
||||||
|
)
|
||||||
|
self.v = torch.zeros(
|
||||||
|
n_layers,
|
||||||
|
max_batch_size,
|
||||||
|
max_seq_len,
|
||||||
|
n_kv_heads,
|
||||||
|
head_dim,
|
||||||
|
device=device,
|
||||||
|
dtype=dtype,
|
||||||
|
)
|
||||||
|
self._slot_len: Dict[int, int] = {}
|
||||||
|
self._task_slot: Dict[str, int] = {}
|
||||||
|
self._free_slots = list(range(max_batch_size))
|
||||||
|
self._device = device
|
||||||
|
|
||||||
|
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
||||||
|
if not self._free_slots:
|
||||||
|
return False
|
||||||
|
slot = self._free_slots.pop(0)
|
||||||
|
self._task_slot[task_id] = slot
|
||||||
|
self._slot_len[slot] = 0
|
||||||
|
return True
|
||||||
|
|
||||||
|
def task_free(self, task_id: str):
|
||||||
|
slot = self._task_slot.pop(task_id, None)
|
||||||
|
if slot is not None:
|
||||||
|
self._slot_len.pop(slot, None)
|
||||||
|
self._free_slots.append(slot)
|
||||||
|
|
||||||
|
def task_extend(self, task_id: str, pos: int) -> bool:
|
||||||
|
return pos < self.max_seq_len
|
||||||
|
|
||||||
|
def bind_tasks(
|
||||||
|
self, task_ids: List[str], total_len: int, device: torch.device
|
||||||
|
) -> ContiguousCacheView:
|
||||||
|
slots = [self._task_slot[tid] for tid in task_ids]
|
||||||
|
batch_indices = torch.tensor(slots, dtype=torch.long, device=device)
|
||||||
|
return ContiguousCacheView(self, batch_indices, total_len)
|
||||||
|
|||||||
@@ -19,13 +19,13 @@ class Executor:
|
|||||||
self,
|
self,
|
||||||
model: AutoModel,
|
model: AutoModel,
|
||||||
tokenizer: AutoTokenizer,
|
tokenizer: AutoTokenizer,
|
||||||
page_cache: KVCache,
|
kv_cache: KVCache,
|
||||||
device: Optional[str] = None,
|
device: Optional[str] = None,
|
||||||
dtype: Optional[torch.dtype] = None,
|
dtype: Optional[torch.dtype] = None,
|
||||||
):
|
):
|
||||||
self.model = model
|
self.model = model
|
||||||
self.tokenizer = tokenizer
|
self.tokenizer = tokenizer
|
||||||
self.page_cache = page_cache
|
self.kv_cache = kv_cache
|
||||||
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
|
||||||
|
|
||||||
@@ -43,7 +43,6 @@ class Executor:
|
|||||||
)
|
)
|
||||||
|
|
||||||
task_ids = [t.task_id for t in tasks]
|
task_ids = [t.task_id for t in tasks]
|
||||||
page_tables = self.page_cache.make_table_tensor(task_ids, self.device)
|
|
||||||
|
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
self.model(
|
self.model(
|
||||||
@@ -53,7 +52,7 @@ class Executor:
|
|||||||
)
|
)
|
||||||
.unsqueeze(0)
|
.unsqueeze(0)
|
||||||
.expand(batch_sz, -1),
|
.expand(batch_sz, -1),
|
||||||
paged_cache=self.page_cache.bind(page_tables, total_len=prompt_len),
|
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]) -> List[int]:
|
||||||
@@ -72,7 +71,6 @@ class Executor:
|
|||||||
total_len = position_ids.max().item() + 1
|
total_len = position_ids.max().item() + 1
|
||||||
|
|
||||||
task_ids = [t.task_id for t in tasks]
|
task_ids = [t.task_id for t in tasks]
|
||||||
page_tables = self.page_cache.make_table_tensor(task_ids, self.device)
|
|
||||||
|
|
||||||
temperatures = torch.tensor([t.temperature for t in tasks], device=self.device)
|
temperatures = torch.tensor([t.temperature for t in tasks], device=self.device)
|
||||||
top_ks = torch.tensor([t.top_k for t in tasks], device=self.device)
|
top_ks = torch.tensor([t.top_k for t in tasks], device=self.device)
|
||||||
@@ -81,7 +79,7 @@ class Executor:
|
|||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
outputs = self.model(
|
outputs = self.model(
|
||||||
input_ids.unsqueeze(1),
|
input_ids.unsqueeze(1),
|
||||||
paged_cache=self.page_cache.bind(page_tables, total_len=total_len),
|
paged_cache=self.kv_cache.bind_tasks(task_ids, total_len, self.device),
|
||||||
position_ids=position_ids.unsqueeze(1),
|
position_ids=position_ids.unsqueeze(1),
|
||||||
)
|
)
|
||||||
logits = outputs["logits"][:, -1, :]
|
logits = outputs["logits"][:, -1, :]
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ from typing import Any, Dict, List, Optional, Tuple
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from astrai.inference.core.cache import KVCache
|
from astrai.inference.core.cache import ContiguousCache, KVCache
|
||||||
from astrai.inference.core.executor import Executor
|
from astrai.inference.core.executor import Executor
|
||||||
from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
|
from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
|
||||||
from astrai.model.automodel import AutoModel
|
from astrai.model.automodel import AutoModel
|
||||||
@@ -14,7 +14,7 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
|
|
||||||
class InferenceScheduler:
|
class InferenceScheduler:
|
||||||
"""Four-phase continuous batching loop: cleanup -> refill -> prefill -> decode."""
|
"""Continuous batching loop: cleanup -> refill -> prefill -> decode (all groups)."""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -23,9 +23,9 @@ class InferenceScheduler:
|
|||||||
max_batch_size: int = 16,
|
max_batch_size: int = 16,
|
||||||
max_seq_len: Optional[int] = None,
|
max_seq_len: Optional[int] = None,
|
||||||
max_prompt_len: int = 2048,
|
max_prompt_len: int = 2048,
|
||||||
page_size: int = 64,
|
|
||||||
device: Optional[str] = None,
|
device: Optional[str] = None,
|
||||||
dtype: Optional[torch.dtype] = None,
|
dtype: Optional[torch.dtype] = None,
|
||||||
|
cache: Optional[KVCache] = None,
|
||||||
):
|
):
|
||||||
config = model.config
|
config = model.config
|
||||||
|
|
||||||
@@ -41,19 +41,20 @@ class InferenceScheduler:
|
|||||||
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
|
||||||
|
|
||||||
n_pages = (
|
head_dim = config.dim // config.n_heads
|
||||||
max_batch_size * (self.max_seq_len + page_size) + page_size - 1
|
|
||||||
) // page_size
|
|
||||||
|
|
||||||
self._page_cache = KVCache(
|
if cache is not None:
|
||||||
config.n_layers,
|
self._cache = cache
|
||||||
n_pages,
|
else:
|
||||||
page_size,
|
self._cache = ContiguousCache(
|
||||||
config.n_kv_heads,
|
config.n_layers,
|
||||||
config.dim // config.n_heads,
|
max_batch_size,
|
||||||
self.device,
|
self.max_seq_len,
|
||||||
self.dtype,
|
config.n_kv_heads,
|
||||||
)
|
head_dim,
|
||||||
|
self.device,
|
||||||
|
self.dtype,
|
||||||
|
)
|
||||||
|
|
||||||
self._task_mgr = TaskManager(
|
self._task_mgr = TaskManager(
|
||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
@@ -65,7 +66,7 @@ class InferenceScheduler:
|
|||||||
self._executor = Executor(
|
self._executor = Executor(
|
||||||
model=model,
|
model=model,
|
||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
page_cache=self._page_cache,
|
kv_cache=self._cache,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
dtype=self.dtype,
|
dtype=self.dtype,
|
||||||
)
|
)
|
||||||
@@ -78,18 +79,19 @@ class InferenceScheduler:
|
|||||||
|
|
||||||
def remove_task(self, task_id: str):
|
def remove_task(self, task_id: str):
|
||||||
for task in self._task_mgr.remove_task(task_id):
|
for task in self._task_mgr.remove_task(task_id):
|
||||||
self._page_cache.task_free(task.task_id)
|
self._cache.task_free(task.task_id)
|
||||||
|
|
||||||
def get_stats(self) -> Dict[str, Any]:
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
return self._task_mgr.get_stats()
|
return self._task_mgr.get_stats()
|
||||||
|
|
||||||
def _run_generation_loop(self):
|
def _run_generation_loop(self):
|
||||||
stop_ids = self._task_mgr.tokenizer.stop_ids
|
stop_ids = self._task_mgr.tokenizer.stop_ids
|
||||||
|
cache = self._cache
|
||||||
try:
|
try:
|
||||||
while not self._stop_event.is_set():
|
while not self._stop_event.is_set():
|
||||||
finished = self._task_mgr.remove_finished_tasks(stop_ids)
|
finished = self._task_mgr.remove_finished_tasks(stop_ids)
|
||||||
for task in finished:
|
for task in finished:
|
||||||
self._page_cache.task_free(task.task_id)
|
cache.task_free(task.task_id)
|
||||||
|
|
||||||
active = self._task_mgr.get_active_tasks()
|
active = self._task_mgr.get_active_tasks()
|
||||||
available = self._task_mgr.max_batch_size - len(active)
|
available = self._task_mgr.max_batch_size - len(active)
|
||||||
@@ -97,7 +99,7 @@ class InferenceScheduler:
|
|||||||
candidates = self._task_mgr.pull_candidates(available)
|
candidates = self._task_mgr.pull_candidates(available)
|
||||||
failed = []
|
failed = []
|
||||||
for task in candidates:
|
for task in candidates:
|
||||||
if self._page_cache.task_alloc(task.task_id, task.prompt_ids):
|
if cache.task_alloc(task.task_id, task.prompt_ids):
|
||||||
self._task_mgr.activate(task)
|
self._task_mgr.activate(task)
|
||||||
else:
|
else:
|
||||||
failed.append(task)
|
failed.append(task)
|
||||||
@@ -112,7 +114,7 @@ class InferenceScheduler:
|
|||||||
t
|
t
|
||||||
for t in self._task_mgr.get_active_tasks()
|
for t in self._task_mgr.get_active_tasks()
|
||||||
if t.output_tokens == 0
|
if t.output_tokens == 0
|
||||||
and self._page_cache.task_cached(t.task_id) < len(t.prompt_ids)
|
and cache.task_cached(t.task_id) < len(t.prompt_ids)
|
||||||
]
|
]
|
||||||
if to_prefill:
|
if to_prefill:
|
||||||
for t in to_prefill:
|
for t in to_prefill:
|
||||||
@@ -122,36 +124,34 @@ class InferenceScheduler:
|
|||||||
for t in to_prefill:
|
for t in to_prefill:
|
||||||
key = (
|
key = (
|
||||||
len(t.prompt_ids),
|
len(t.prompt_ids),
|
||||||
self._page_cache.task_cached(t.task_id),
|
cache.task_cached(t.task_id),
|
||||||
)
|
)
|
||||||
groups.setdefault(key, []).append(t)
|
groups.setdefault(key, []).append(t)
|
||||||
|
|
||||||
for (prompt_len, start_pos), group in groups.items():
|
for (prompt_len, start_pos), group in groups.items():
|
||||||
self._executor.execute_prefill(group, prompt_len, start_pos)
|
self._executor.execute_prefill(group, prompt_len, start_pos)
|
||||||
start_logical_page = start_pos // self._page_cache.page_size
|
start_logical_page = start_pos // getattr(
|
||||||
|
cache, "page_size", 64
|
||||||
|
)
|
||||||
for t in group:
|
for t in group:
|
||||||
self._page_cache.task_record_hashes(
|
cache.task_record_hashes(
|
||||||
t.task_id,
|
t.task_id, t.prompt_ids, start_logical_page
|
||||||
t.prompt_ids,
|
|
||||||
start_logical_page=start_logical_page,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
pos_groups: Dict[int, List[Task]] = {}
|
pos_groups: Dict[int, List[Task]] = {}
|
||||||
for t in self._task_mgr.get_active_tasks():
|
for t in self._task_mgr.get_active_tasks():
|
||||||
pos_groups.setdefault(t.next_pos, []).append(t)
|
pos_groups.setdefault(t.next_pos, []).append(t)
|
||||||
|
|
||||||
if pos_groups:
|
for next_pos in sorted(pos_groups.keys()):
|
||||||
best_key = max(pos_groups, key=lambda k: len(pos_groups[k]))
|
group = sorted(pos_groups[next_pos], key=lambda t: t.task_id)
|
||||||
group = sorted(pos_groups[best_key], key=lambda t: t.task_id)
|
|
||||||
|
|
||||||
valid: List[Task] = []
|
valid: List[Task] = []
|
||||||
for t in group:
|
for t in group:
|
||||||
if self._page_cache.task_extend(t.task_id, t.next_pos):
|
if cache.task_extend(t.task_id, t.next_pos):
|
||||||
valid.append(t)
|
valid.append(t)
|
||||||
else:
|
else:
|
||||||
t.status = TaskStatus.ABORTED
|
t.status = TaskStatus.ABORTED
|
||||||
if t.stream_callback:
|
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||||
t.stream_callback(STOP)
|
|
||||||
|
|
||||||
if valid:
|
if valid:
|
||||||
next_tokens = self._executor.execute_decode(valid)
|
next_tokens = self._executor.execute_decode(valid)
|
||||||
@@ -159,32 +159,23 @@ class InferenceScheduler:
|
|||||||
for t, ntok in zip(valid, next_tokens):
|
for t, ntok in zip(valid, next_tokens):
|
||||||
t.output_ids.append(ntok)
|
t.output_ids.append(ntok)
|
||||||
t.output_tokens += 1
|
t.output_tokens += 1
|
||||||
pos = t.input_tokens + t.output_tokens
|
self._task_mgr.invoke_callback(
|
||||||
extend_ok = self._page_cache.task_extend(t.task_id, pos)
|
t.task_id,
|
||||||
if t.stream_callback:
|
self._task_mgr.tokenizer.decode([ntok]),
|
||||||
t.stream_callback(
|
)
|
||||||
self._task_mgr.tokenizer.decode([ntok])
|
|
||||||
)
|
|
||||||
if not extend_ok:
|
|
||||||
t.status = TaskStatus.ABORTED
|
|
||||||
if t.stream_callback:
|
|
||||||
t.stream_callback(STOP)
|
|
||||||
|
|
||||||
for t in valid:
|
for t in valid:
|
||||||
if t.is_finished(stop_ids):
|
if t.is_finished(stop_ids):
|
||||||
if t.stream_callback:
|
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||||
t.stream_callback(STOP)
|
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self._stop_event.set()
|
self._stop_event.set()
|
||||||
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
|
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
|
||||||
for task in self._task_mgr.get_active_tasks():
|
for task in self._task_mgr.get_active_tasks():
|
||||||
if task.stream_callback:
|
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||||
task.stream_callback(STOP)
|
cache.task_free(task.task_id)
|
||||||
self._page_cache.task_free(task.task_id)
|
|
||||||
for task in self._task_mgr.get_waiting_tasks():
|
for task in self._task_mgr.get_waiting_tasks():
|
||||||
if task.stream_callback:
|
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||||
task.stream_callback(STOP)
|
|
||||||
self._task_mgr.clear_queues()
|
self._task_mgr.clear_queues()
|
||||||
|
|
||||||
def start(self):
|
def start(self):
|
||||||
@@ -202,12 +193,10 @@ class InferenceScheduler:
|
|||||||
self._loop_thread.join(timeout=2.0)
|
self._loop_thread.join(timeout=2.0)
|
||||||
self._loop_thread = None
|
self._loop_thread = None
|
||||||
for task in self._task_mgr.get_active_tasks():
|
for task in self._task_mgr.get_active_tasks():
|
||||||
if task.stream_callback:
|
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||||
task.stream_callback(STOP)
|
self._cache.task_free(task.task_id)
|
||||||
self._page_cache.task_free(task.task_id)
|
|
||||||
for task in self._task_mgr.get_waiting_tasks():
|
for task in self._task_mgr.get_waiting_tasks():
|
||||||
if task.stream_callback:
|
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||||
task.stream_callback(STOP)
|
|
||||||
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()
|
||||||
|
|||||||
@@ -33,7 +33,6 @@ class Task:
|
|||||||
temperature: float = 1.0,
|
temperature: float = 1.0,
|
||||||
top_p: float = 1.0,
|
top_p: float = 1.0,
|
||||||
top_k: int = 50,
|
top_k: int = 50,
|
||||||
stream_callback: Optional[Callable[[str], None]] = None,
|
|
||||||
):
|
):
|
||||||
self.task_id = task_id
|
self.task_id = task_id
|
||||||
self.prompt_ids = prompt_ids
|
self.prompt_ids = prompt_ids
|
||||||
@@ -48,7 +47,6 @@ class Task:
|
|||||||
self.output_tokens: int = 0
|
self.output_tokens: int = 0
|
||||||
self.arrival_time = time.time()
|
self.arrival_time = time.time()
|
||||||
self.finish_time: Optional[float] = None
|
self.finish_time: Optional[float] = None
|
||||||
self.stream_callback = stream_callback
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def next_pos(self) -> int:
|
def next_pos(self) -> int:
|
||||||
@@ -79,6 +77,7 @@ class TaskManager:
|
|||||||
|
|
||||||
self.waiting_queue: Deque[Task] = deque()
|
self.waiting_queue: Deque[Task] = deque()
|
||||||
self.active_tasks: List[Task] = []
|
self.active_tasks: List[Task] = []
|
||||||
|
self._callbacks: Dict[str, Callable[[str], None]] = {}
|
||||||
|
|
||||||
self._task_event = threading.Event()
|
self._task_event = threading.Event()
|
||||||
self._lock = threading.Lock()
|
self._lock = threading.Lock()
|
||||||
@@ -117,12 +116,13 @@ class TaskManager:
|
|||||||
temperature=temperature,
|
temperature=temperature,
|
||||||
top_p=top_p,
|
top_p=top_p,
|
||||||
top_k=top_k,
|
top_k=top_k,
|
||||||
stream_callback=stream_callback,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self.waiting_queue.append(task)
|
self.waiting_queue.append(task)
|
||||||
self._total_tasks += 1
|
self._total_tasks += 1
|
||||||
|
if stream_callback:
|
||||||
|
self._callbacks[task_id] = stream_callback
|
||||||
|
|
||||||
self._task_event.set()
|
self._task_event.set()
|
||||||
return task_id
|
return task_id
|
||||||
@@ -134,8 +134,14 @@ class TaskManager:
|
|||||||
t for t in self.waiting_queue if t.task_id != task_id
|
t for t in self.waiting_queue if t.task_id != task_id
|
||||||
)
|
)
|
||||||
self.active_tasks = [t for t in self.active_tasks if t.task_id != task_id]
|
self.active_tasks = [t for t in self.active_tasks if t.task_id != task_id]
|
||||||
|
self._callbacks.pop(task_id, None)
|
||||||
return removed_active
|
return removed_active
|
||||||
|
|
||||||
|
def invoke_callback(self, task_id: str, token: str):
|
||||||
|
cb = self._callbacks.get(task_id)
|
||||||
|
if cb:
|
||||||
|
cb(token)
|
||||||
|
|
||||||
def get_stats(self) -> Dict[str, Any]:
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
return {
|
return {
|
||||||
"total_tasks": self._total_tasks,
|
"total_tasks": self._total_tasks,
|
||||||
@@ -204,6 +210,7 @@ class TaskManager:
|
|||||||
with self._lock:
|
with self._lock:
|
||||||
self.waiting_queue.clear()
|
self.waiting_queue.clear()
|
||||||
self.active_tasks.clear()
|
self.active_tasks.clear()
|
||||||
|
self._callbacks.clear()
|
||||||
|
|
||||||
def wake(self):
|
def wake(self):
|
||||||
self._task_event.set()
|
self._task_event.set()
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Tuple,
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|
||||||
|
from astrai.inference.core.cache import KVCache
|
||||||
from astrai.inference.core.scheduler import InferenceScheduler
|
from astrai.inference.core.scheduler import InferenceScheduler
|
||||||
from astrai.inference.core.task import STOP
|
from astrai.inference.core.task import STOP
|
||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
@@ -101,6 +102,7 @@ class InferenceEngine:
|
|||||||
max_seq_len: Optional[int] = None,
|
max_seq_len: Optional[int] = None,
|
||||||
max_prompt_len: int = 2048,
|
max_prompt_len: int = 2048,
|
||||||
page_size: int = 128,
|
page_size: int = 128,
|
||||||
|
cache: Optional[KVCache] = None,
|
||||||
):
|
):
|
||||||
self.model = model
|
self.model = model
|
||||||
self.tokenizer = tokenizer
|
self.tokenizer = tokenizer
|
||||||
@@ -110,7 +112,7 @@ class InferenceEngine:
|
|||||||
max_batch_size=max_batch_size,
|
max_batch_size=max_batch_size,
|
||||||
max_seq_len=max_seq_len,
|
max_seq_len=max_seq_len,
|
||||||
max_prompt_len=max_prompt_len,
|
max_prompt_len=max_prompt_len,
|
||||||
page_size=page_size,
|
cache=cache,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.scheduler.start()
|
self.scheduler.start()
|
||||||
|
|||||||
@@ -6,7 +6,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.inference.core.cache import KvcacheView
|
from astrai.inference.core.cache import CacheView
|
||||||
from astrai.model.components.linear import Linear
|
from astrai.model.components.linear import Linear
|
||||||
from astrai.model.components.norm import RMSNorm
|
from astrai.model.components.norm import RMSNorm
|
||||||
from astrai.model.components.rope import apply_rotary_emb
|
from astrai.model.components.rope import apply_rotary_emb
|
||||||
@@ -75,7 +75,7 @@ class GQA(nn.Module):
|
|||||||
x: Tensor,
|
x: Tensor,
|
||||||
rotary_emb: Tensor,
|
rotary_emb: Tensor,
|
||||||
attn_mask: Tensor = None,
|
attn_mask: Tensor = None,
|
||||||
paged_cache: Optional[KvcacheView] = None,
|
paged_cache: Optional[CacheView] = None,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
is_causal = attn_mask is None
|
is_causal = attn_mask is None
|
||||||
|
|
||||||
@@ -162,7 +162,7 @@ class MLA(nn.Module):
|
|||||||
x: Tensor,
|
x: Tensor,
|
||||||
rotary_emb: Tensor,
|
rotary_emb: Tensor,
|
||||||
attn_mask: Tensor = None,
|
attn_mask: Tensor = None,
|
||||||
paged_cache: Optional[KvcacheView] = None,
|
paged_cache: Optional[CacheView] = None,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
bsz, seq_len, _ = x.size()
|
bsz, seq_len, _ = x.size()
|
||||||
is_causal = attn_mask is None
|
is_causal = attn_mask is None
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ from typing import Optional
|
|||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.inference.core.cache import KvcacheView
|
from astrai.inference.core.cache import CacheView
|
||||||
from astrai.model.components.attention import AttnFactory
|
from astrai.model.components.attention import AttnFactory
|
||||||
from astrai.model.components.mlp import FFNFactory
|
from astrai.model.components.mlp import FFNFactory
|
||||||
from astrai.model.components.norm import RMSNorm
|
from astrai.model.components.norm import RMSNorm
|
||||||
@@ -25,7 +25,7 @@ class DecoderBlock(nn.Module):
|
|||||||
x: Tensor,
|
x: Tensor,
|
||||||
rotary_emb: Tensor,
|
rotary_emb: Tensor,
|
||||||
attention_mask: Optional[Tensor] = None,
|
attention_mask: Optional[Tensor] = None,
|
||||||
paged_cache: Optional[KvcacheView] = None,
|
paged_cache: Optional[CacheView] = None,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
attn_output = self.attention(
|
attn_output = self.attention(
|
||||||
self.input_norm(x),
|
self.input_norm(x),
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ import torch.nn as nn
|
|||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
from astrai.inference.core.cache import KvcacheView
|
from astrai.inference.core.cache import CacheView
|
||||||
from astrai.model.automodel import AutoModel
|
from astrai.model.automodel import AutoModel
|
||||||
from astrai.model.components.decoder_block import DecoderBlock
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
from astrai.model.components.embedding import Embedding
|
from astrai.model.components.embedding import Embedding
|
||||||
@@ -112,7 +112,7 @@ class AutoRegressiveLM(AutoModel):
|
|||||||
self,
|
self,
|
||||||
input_ids: Tensor,
|
input_ids: Tensor,
|
||||||
input_mask: Optional[Tensor] = None,
|
input_mask: Optional[Tensor] = None,
|
||||||
paged_cache: Optional[KvcacheView] = None,
|
paged_cache: Optional[CacheView] = None,
|
||||||
position_ids: Optional[Tensor] = None,
|
position_ids: Optional[Tensor] = None,
|
||||||
) -> Dict[str, Tensor]:
|
) -> Dict[str, Tensor]:
|
||||||
assert input_ids.ndim == 2
|
assert input_ids.ndim == 2
|
||||||
|
|||||||
@@ -1,14 +1,21 @@
|
|||||||
import os
|
import os
|
||||||
|
import socket
|
||||||
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
|
||||||
from typing import Callable
|
from typing import Callable, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.distributed as dist
|
import torch.distributed as dist
|
||||||
import torch.multiprocessing as mp
|
import torch.multiprocessing as mp
|
||||||
|
|
||||||
|
|
||||||
|
def find_free_port() -> str:
|
||||||
|
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||||
|
s.bind(("", 0))
|
||||||
|
return str(s.getsockname()[1])
|
||||||
|
|
||||||
|
|
||||||
def get_current_device():
|
def get_current_device():
|
||||||
return os.environ["LOCAL_DEVICE"]
|
return os.environ["LOCAL_DEVICE"]
|
||||||
|
|
||||||
@@ -217,11 +224,13 @@ def spawn_parallel_fn(
|
|||||||
world_size: int,
|
world_size: int,
|
||||||
backend: str = "nccl",
|
backend: str = "nccl",
|
||||||
master_addr: str = "localhost",
|
master_addr: str = "localhost",
|
||||||
master_port: str = "29500",
|
master_port: Optional[str] = None,
|
||||||
device_type: str = "cuda",
|
device_type: str = "cuda",
|
||||||
start_method: str = "spawn",
|
start_method: str = "spawn",
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
|
if master_port is None:
|
||||||
|
master_port = find_free_port()
|
||||||
launcher = _detect_launcher()
|
launcher = _detect_launcher()
|
||||||
if launcher in ("torchelastic", "torchrun", "external"):
|
if launcher in ("torchelastic", "torchrun", "external"):
|
||||||
strategy = TorchrunStrategy(
|
strategy = TorchrunStrategy(
|
||||||
|
|||||||
@@ -1,7 +1,9 @@
|
|||||||
from astrai.preprocessing.builder import (
|
from astrai.preprocessing.builder import (
|
||||||
BaseMaskBuilder,
|
BaseMaskBuilder,
|
||||||
MaskBuilderFactory,
|
MaskBuilderFactory,
|
||||||
|
MultiOutputMaskBuilder,
|
||||||
SectionedMaskBuilder,
|
SectionedMaskBuilder,
|
||||||
|
SingleOutputMaskBuilder,
|
||||||
)
|
)
|
||||||
from astrai.preprocessing.packing import (
|
from astrai.preprocessing.packing import (
|
||||||
PackingStrategy,
|
PackingStrategy,
|
||||||
@@ -20,12 +22,14 @@ from astrai.preprocessing.writer import (
|
|||||||
__all__ = [
|
__all__ = [
|
||||||
"BaseMaskBuilder",
|
"BaseMaskBuilder",
|
||||||
"MaskBuilderFactory",
|
"MaskBuilderFactory",
|
||||||
|
"MultiOutputMaskBuilder",
|
||||||
"PackingStrategy",
|
"PackingStrategy",
|
||||||
"PackingStrategyFactory",
|
"PackingStrategyFactory",
|
||||||
"Pipeline",
|
"Pipeline",
|
||||||
"PositionIdStrategy",
|
"PositionIdStrategy",
|
||||||
"PositionIdStrategyFactory",
|
"PositionIdStrategyFactory",
|
||||||
"SectionedMaskBuilder",
|
"SectionedMaskBuilder",
|
||||||
|
"SingleOutputMaskBuilder",
|
||||||
"StoreWriter",
|
"StoreWriter",
|
||||||
"StoreWriterFactory",
|
"StoreWriterFactory",
|
||||||
"filter_by_length",
|
"filter_by_length",
|
||||||
|
|||||||
@@ -1,8 +1,10 @@
|
|||||||
"""Mask building for preprocessing pipeline.
|
"""Mask building for preprocessing pipeline.
|
||||||
|
|
||||||
:class:`SectionRenderer` converts section specs into token ids and loss
|
:class:`SectionRenderer` converts section specs into token ids and loss
|
||||||
masks (template / text / value extraction). :class:`SectionedMaskBuilder`
|
masks (template / text / value extraction). :class:`SingleOutputMaskBuilder`
|
||||||
orchestrates single-output / multi-output (DPO / GRPO) assembly.
|
handles single-output (SFT / pretrain), :class:`MultiOutputMaskBuilder`
|
||||||
|
handles multi-output (DPO / GRPO), and :class:`SectionedMaskBuilder`
|
||||||
|
orchestrates both modes as a façade.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
@@ -212,42 +214,17 @@ class MaskBuilderFactory(BaseFactory["BaseMaskBuilder"]):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
@MaskBuilderFactory.register("sectioned")
|
@MaskBuilderFactory.register("single")
|
||||||
class SectionedMaskBuilder(BaseMaskBuilder):
|
class SingleOutputMaskBuilder(BaseMaskBuilder):
|
||||||
"""Config-driven builder supporting single and multi-output modes.
|
"""Build a single output sequence with optional loss mask.
|
||||||
|
|
||||||
Single-output::
|
Expects ``config.input.sections`` (list of section specs).
|
||||||
|
|
||||||
{"input": {"sections": [
|
|
||||||
{"field": "messages", "action": "$role", "template": true}
|
|
||||||
]}}
|
|
||||||
→ {"sequence": [...], "loss_mask": [...], "domain": "..."}
|
|
||||||
|
|
||||||
Multi-output (DPO / GRPO)::
|
|
||||||
|
|
||||||
{"input": {"sources": {
|
|
||||||
"chosen": {"sections": [{"field": "chosen", "action": "$role", "template": true}]},
|
|
||||||
"rejected": {"sections": [{"field": "rejected", "action": "$role", "template": true}]},
|
|
||||||
}}}
|
|
||||||
→ {"chosen": [...], "chosen_mask": [...], "rejected": [...], "rejected_mask": [...], "domain": "..."}
|
|
||||||
|
|
||||||
Output spec fields::
|
|
||||||
|
|
||||||
sections – list of section specs (same format as single-output)
|
|
||||||
list_field – True when JSONL field holds a list (GRPO responses)
|
|
||||||
mask_key – explicit loss-mask output key (default: ``"{output_key}_mask"``)
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self):
|
def __init__(self, renderer: Optional[SectionRenderer] = None):
|
||||||
self.renderer = SectionRenderer()
|
self.renderer = renderer or SectionRenderer()
|
||||||
|
|
||||||
def build(self, item: dict, config, tokenizer) -> Optional[dict]:
|
def build(self, item: dict, config, tokenizer) -> Optional[dict]:
|
||||||
sources_spec = getattr(config.input, "sources", None)
|
|
||||||
if sources_spec:
|
|
||||||
return self._build_multi(item, sources_spec, config, tokenizer)
|
|
||||||
return self._build_single(item, config, tokenizer)
|
|
||||||
|
|
||||||
def _build_single(self, item: dict, config, tokenizer) -> Optional[dict]:
|
|
||||||
sections = config.input.sections
|
sections = config.input.sections
|
||||||
if not sections:
|
if not sections:
|
||||||
return None
|
return None
|
||||||
@@ -266,9 +243,22 @@ class SectionedMaskBuilder(BaseMaskBuilder):
|
|||||||
result["loss_mask"] = mask
|
result["loss_mask"] = mask
|
||||||
return result
|
return result
|
||||||
|
|
||||||
def _build_multi(
|
|
||||||
self, item: dict, sources_spec: dict, config, tokenizer
|
@MaskBuilderFactory.register("multi")
|
||||||
) -> Optional[dict]:
|
class MultiOutputMaskBuilder(BaseMaskBuilder):
|
||||||
|
"""Build multiple output sequences (DPO / GRPO).
|
||||||
|
|
||||||
|
Expects ``config.input.sources`` (dict of output_key → spec).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, renderer: Optional[SectionRenderer] = None):
|
||||||
|
self.renderer = renderer or SectionRenderer()
|
||||||
|
|
||||||
|
def build(self, item: dict, config, tokenizer) -> Optional[dict]:
|
||||||
|
sources_spec = getattr(config.input, "sources", None)
|
||||||
|
if not sources_spec:
|
||||||
|
return None
|
||||||
|
|
||||||
result: dict = {}
|
result: dict = {}
|
||||||
any_output = False
|
any_output = False
|
||||||
|
|
||||||
@@ -313,3 +303,22 @@ class SectionedMaskBuilder(BaseMaskBuilder):
|
|||||||
|
|
||||||
result["domain"] = _extract_domain(item, config.output.domain_key)
|
result["domain"] = _extract_domain(item, config.output.domain_key)
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
@MaskBuilderFactory.register("sectioned")
|
||||||
|
class SectionedMaskBuilder(BaseMaskBuilder):
|
||||||
|
"""Façade that dispatches to SingleOutputMaskBuilder or MultiOutputMaskBuilder.
|
||||||
|
|
||||||
|
Preserves backward compatibility for existing configs and code that rely
|
||||||
|
on the ``"sectioned"`` factory name.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self._single = SingleOutputMaskBuilder()
|
||||||
|
self._multi = MultiOutputMaskBuilder()
|
||||||
|
|
||||||
|
def build(self, item: dict, config, tokenizer) -> Optional[dict]:
|
||||||
|
sources_spec = getattr(config.input, "sources", None)
|
||||||
|
if sources_spec:
|
||||||
|
return self._multi.build(item, config, tokenizer)
|
||||||
|
return self._single.build(item, config, tokenizer)
|
||||||
|
|||||||
@@ -144,7 +144,7 @@ class Pipeline:
|
|||||||
for key in list(bucket.keys()):
|
for key in list(bucket.keys()):
|
||||||
if key in result:
|
if key in result:
|
||||||
continue
|
continue
|
||||||
bucket[key].append([1] * len(ids))
|
bucket[key].append([0] * len(ids))
|
||||||
|
|
||||||
def _iter_items(self):
|
def _iter_items(self):
|
||||||
for path in self.paths:
|
for path in self.paths:
|
||||||
@@ -160,6 +160,12 @@ class Pipeline:
|
|||||||
idx = shard_idx[domain]
|
idx = shard_idx[domain]
|
||||||
|
|
||||||
pp = self.config.preprocessing
|
pp = self.config.preprocessing
|
||||||
|
original_sequences = keys.get("sequence", [])
|
||||||
|
mode = self.config.output.position_ids_mode
|
||||||
|
|
||||||
|
if mode == "doc_reset" and original_sequences:
|
||||||
|
keys["position_ids"] = [list(range(len(s))) for s in original_sequences]
|
||||||
|
|
||||||
keys = self._packer.apply(dict(keys), pp.max_packed_len, pp.truncation_mode)
|
keys = self._packer.apply(dict(keys), pp.max_packed_len, pp.truncation_mode)
|
||||||
|
|
||||||
tensors: Dict[str, List[torch.Tensor]] = {}
|
tensors: Dict[str, List[torch.Tensor]] = {}
|
||||||
@@ -171,9 +177,10 @@ class Pipeline:
|
|||||||
torch.tensor(list(chain.from_iterable(ids_list)), dtype=dt)
|
torch.tensor(list(chain.from_iterable(ids_list)), dtype=dt)
|
||||||
]
|
]
|
||||||
|
|
||||||
pos_ids = self._position_id.generate(keys.get("sequence", []))
|
if mode == "continuous" and original_sequences:
|
||||||
if pos_ids:
|
pos_ids = self._position_id.generate(keys.get("sequence", []))
|
||||||
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
|
if pos_ids:
|
||||||
|
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
|
||||||
|
|
||||||
self._writer.save(self.output_dir, domain, idx, tensors)
|
self._writer.save(self.output_dir, domain, idx, tensors)
|
||||||
shard_idx[domain] = idx + 1
|
shard_idx[domain] = idx + 1
|
||||||
|
|||||||
@@ -26,7 +26,10 @@ def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
|
|||||||
tensor_group: Dict[str, List[Tensor]] = {}
|
tensor_group: Dict[str, List[Tensor]] = {}
|
||||||
|
|
||||||
root_path = Path(file_path)
|
root_path = Path(file_path)
|
||||||
h5_files = list(root_path.rglob("*.h5")) + list(root_path.rglob("*.hdf5"))
|
if root_path.is_file() and root_path.suffix in (".h5", ".hdf5"):
|
||||||
|
h5_files = [root_path]
|
||||||
|
else:
|
||||||
|
h5_files = list(root_path.rglob("*.h5")) + list(root_path.rglob("*.hdf5"))
|
||||||
|
|
||||||
for h5_file in h5_files:
|
for h5_file in h5_files:
|
||||||
with h5py.File(h5_file, "r") as f:
|
with h5py.File(h5_file, "r") as f:
|
||||||
|
|||||||
@@ -209,7 +209,6 @@ class ProgressBarCallback(TrainCallback):
|
|||||||
|
|
||||||
@only_on_rank(0)
|
@only_on_rank(0)
|
||||||
def on_optimizer_step(self, context: TrainContext):
|
def on_optimizer_step(self, context: TrainContext):
|
||||||
self.progress_bar.update(1)
|
|
||||||
postfix = {
|
postfix = {
|
||||||
"step": context.optimizer_step,
|
"step": context.optimizer_step,
|
||||||
"loss": f"{context.loss:.4f}",
|
"loss": f"{context.loss:.4f}",
|
||||||
@@ -220,6 +219,7 @@ class ProgressBarCallback(TrainCallback):
|
|||||||
if context.val_loss is not None:
|
if context.val_loss is not None:
|
||||||
postfix["val_loss"] = f"{context.val_loss:.4f}"
|
postfix["val_loss"] = f"{context.val_loss:.4f}"
|
||||||
self.progress_bar.set_postfix(postfix)
|
self.progress_bar.set_postfix(postfix)
|
||||||
|
self.progress_bar.update(1)
|
||||||
|
|
||||||
@only_on_rank(0)
|
@only_on_rank(0)
|
||||||
def on_epoch_end(self, context: TrainContext):
|
def on_epoch_end(self, context: TrainContext):
|
||||||
|
|||||||
@@ -0,0 +1,2 @@
|
|||||||
|
# Source directory for CUDA kernels — build-time only.
|
||||||
|
# Compiled .so files live in astrAI/_ext/.
|
||||||
@@ -0,0 +1,48 @@
|
|||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
def _arch_flags() -> list[str]:
|
||||||
|
import torch
|
||||||
|
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
cap = torch.cuda.get_device_capability()
|
||||||
|
else:
|
||||||
|
cap = (8, 0)
|
||||||
|
ver = f"{cap[0]}{cap[1]}"
|
||||||
|
flags = [f"-gencode=arch=compute_{ver},code=sm_{ver}"]
|
||||||
|
# tensor-core mma path (mma.sync.m16n8k16.bf16) requires sm_80+; decide the
|
||||||
|
# kernel dispatch at build time via this define rather than at runtime.
|
||||||
|
if cap[0] < 8:
|
||||||
|
flags.append("-DASTRAI_NO_MMA")
|
||||||
|
return flags
|
||||||
|
|
||||||
|
|
||||||
|
_kernels_dir = Path("csrc/kernels")
|
||||||
|
REGISTRY: dict[str, dict] = {}
|
||||||
|
|
||||||
|
CXX_FLAGS = ["-O3", "-funroll-loops"]
|
||||||
|
NVCC_FLAGS = [
|
||||||
|
"-O3",
|
||||||
|
"--expt-relaxed-constexpr",
|
||||||
|
"--use_fast_math",
|
||||||
|
"--ptxas-options=-O3,-v",
|
||||||
|
"--extra-device-vectorization",
|
||||||
|
"--threads=8",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def register(name: str, sources: list[str] | None = None, **kwargs):
|
||||||
|
if sources is None:
|
||||||
|
sources = [str(_kernels_dir / f"{name}.cu")]
|
||||||
|
REGISTRY[name] = {
|
||||||
|
"sources": sources,
|
||||||
|
"cxx_flags": [*CXX_FLAGS],
|
||||||
|
"nvcc_flags": [*NVCC_FLAGS, *_arch_flags()],
|
||||||
|
"extra_link_args": kwargs.pop("extra_link_args", []),
|
||||||
|
**kwargs,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
register("attn_decode")
|
||||||
|
register("attn_prefill")
|
||||||
|
register("attn_paged_decode")
|
||||||
@@ -0,0 +1,54 @@
|
|||||||
|
#pragma once
|
||||||
|
|
||||||
|
|
||||||
|
template<typename T, typename AT = float>
|
||||||
|
struct AttentionParams {
|
||||||
|
int batch;
|
||||||
|
int q_head;
|
||||||
|
int kv_head;
|
||||||
|
int q_len;
|
||||||
|
int kv_len;
|
||||||
|
int head_dim;
|
||||||
|
int use_mask;
|
||||||
|
int is_causal;
|
||||||
|
int causal_offset;
|
||||||
|
int num_splits;
|
||||||
|
float scale;
|
||||||
|
|
||||||
|
const T* __restrict__ q;
|
||||||
|
const T* __restrict__ k;
|
||||||
|
const T* __restrict__ v;
|
||||||
|
const bool* __restrict__ mask;
|
||||||
|
|
||||||
|
T* __restrict__ o;
|
||||||
|
AT* __restrict__ o_part;
|
||||||
|
AT* __restrict__ ml_part;
|
||||||
|
};
|
||||||
|
|
||||||
|
template<typename T, typename AT = float>
|
||||||
|
struct PagedAttentionParams {
|
||||||
|
int batch;
|
||||||
|
int q_head;
|
||||||
|
int kv_head;
|
||||||
|
int q_len;
|
||||||
|
int kv_len;
|
||||||
|
int head_dim;
|
||||||
|
int use_mask;
|
||||||
|
int is_causal;
|
||||||
|
int causal_offset;
|
||||||
|
float scale;
|
||||||
|
|
||||||
|
int num_splits;
|
||||||
|
int page_size;
|
||||||
|
int max_pages;
|
||||||
|
|
||||||
|
const T* __restrict__ q;
|
||||||
|
const T* __restrict__ k_cache;
|
||||||
|
const T* __restrict__ v_cache;
|
||||||
|
const bool* __restrict__ mask;
|
||||||
|
const int64_t* __restrict__ page_table;
|
||||||
|
|
||||||
|
T* __restrict__ o;
|
||||||
|
AT* __restrict__ o_part;
|
||||||
|
AT* __restrict__ ml_part;
|
||||||
|
};
|
||||||
@@ -0,0 +1,113 @@
|
|||||||
|
#include "attn_decode_split_kv.cuh"
|
||||||
|
#include "attn_entry_utils.cuh"
|
||||||
|
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
#include "attn_decode_split_kv_mma.cuh"
|
||||||
|
#endif
|
||||||
|
|
||||||
|
static int decode_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)));
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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 = decode_num_splits(p.batch * p.kv_head, chunks_total);
|
||||||
|
|
||||||
|
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 ml_part = torch::empty({p.batch, p.q_head, p.num_splits, 2}, fopt);
|
||||||
|
p.o_part = o_part.data_ptr<float>();
|
||||||
|
p.ml_part = ml_part.data_ptr<float>();
|
||||||
|
|
||||||
|
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 = decode_num_splits(p.batch * p.kv_head, tiles_total);
|
||||||
|
|
||||||
|
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 ml_part = torch::empty({p.batch, p.q_head, p.num_splits, 2}, fopt);
|
||||||
|
p.o_part = o_part.data_ptr<float>();
|
||||||
|
p.ml_part = ml_part.data_ptr<float>();
|
||||||
|
|
||||||
|
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 (!p.use_mask && G >= 1 && G <= 16) {
|
||||||
|
launch_mma_decode<HEAD_DIM, 32>(p);
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
#endif
|
||||||
|
launch_scalar_decode(p);
|
||||||
|
}
|
||||||
|
|
||||||
|
torch::Tensor attn_decode(
|
||||||
|
torch::Tensor q,
|
||||||
|
torch::Tensor k,
|
||||||
|
torch::Tensor v,
|
||||||
|
c10::optional<torch::Tensor> mask,
|
||||||
|
bool is_causal = false,
|
||||||
|
int64_t causal_offset = 0,
|
||||||
|
c10::optional<double> scale = c10::nullopt
|
||||||
|
) {
|
||||||
|
AttentionParams<bf16> p;
|
||||||
|
attn_pack_params(q, k, v, mask, is_causal, causal_offset, scale, p);
|
||||||
|
TORCH_CHECK(p.q_len == 1, "Q seq_len must be 1");
|
||||||
|
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
|
||||||
|
|
||||||
|
auto O = torch::empty_like(q);
|
||||||
|
p.o = (bf16*)O.data_ptr();
|
||||||
|
|
||||||
|
switch (p.head_dim) {
|
||||||
|
case 32:
|
||||||
|
dispatch_decode<32>(p);
|
||||||
|
break;
|
||||||
|
case 64:
|
||||||
|
dispatch_decode<64>(p);
|
||||||
|
break;
|
||||||
|
case 128:
|
||||||
|
dispatch_decode<128>(p);
|
||||||
|
break;
|
||||||
|
case 256:
|
||||||
|
dispatch_decode<256>(p);
|
||||||
|
break;
|
||||||
|
default:
|
||||||
|
TORCH_CHECK(false, "decode: unsupported head_dim ", p.head_dim,
|
||||||
|
" (supported: 32, 64, 128, 256)");
|
||||||
|
}
|
||||||
|
return O;
|
||||||
|
}
|
||||||
|
|
||||||
|
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||||
|
m.def("attn_decode", &attn_decode,
|
||||||
|
py::arg("q"),
|
||||||
|
py::arg("k"),
|
||||||
|
py::arg("v"),
|
||||||
|
py::arg("mask") = py::none(),
|
||||||
|
py::arg("is_causal") = false,
|
||||||
|
py::arg("causal_offset") = 0,
|
||||||
|
py::arg("scale") = py::none(),
|
||||||
|
"GQA decode (tensor-core head-packing on sm_80+, scalar fallback)");
|
||||||
|
}
|
||||||
@@ -0,0 +1,116 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include <float.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
|
||||||
|
using bf16 = __nv_bfloat16;
|
||||||
|
constexpr int DC_CHUNK = 64;
|
||||||
|
|
||||||
|
__device__ inline float warp_reduce_sum(float val) {
|
||||||
|
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) {
|
||||||
|
int batch = blockIdx.x / p.kv_head;
|
||||||
|
int kv_head = blockIdx.x % p.kv_head;
|
||||||
|
int split = blockIdx.z;
|
||||||
|
int group_size = blockDim.y;
|
||||||
|
int q_head = kv_head * group_size + threadIdx.y;
|
||||||
|
int lane = threadIdx.x;
|
||||||
|
int hd_per_thread = p.head_dim / 32;
|
||||||
|
|
||||||
|
float q_reg[8];
|
||||||
|
int q_off = ((batch * p.q_head + q_head) * 1) * p.head_dim + lane * hd_per_thread;
|
||||||
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
|
q_reg[i] = __bfloat162float(p.q[q_off + i]);
|
||||||
|
|
||||||
|
int kv_base = ((batch * p.kv_head + kv_head) * p.kv_len) * p.head_dim;
|
||||||
|
int mask_base = batch * p.kv_len;
|
||||||
|
|
||||||
|
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
|
||||||
|
|
||||||
|
extern __shared__ __align__(16) bf16 k_smem[];
|
||||||
|
|
||||||
|
// Split-KV: each split processes a contiguous subset of chunks
|
||||||
|
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
||||||
|
int chunks_per_split = (chunks_total + p.num_splits - 1) / p.num_splits;
|
||||||
|
int ch_begin = split * chunks_per_split;
|
||||||
|
int ch_end = min(chunks_total, ch_begin + chunks_per_split);
|
||||||
|
|
||||||
|
for (int ci = ch_begin; ci < ch_end; ci++) {
|
||||||
|
int chunk_start = ci * DC_CHUNK;
|
||||||
|
int this_chunk = min(DC_CHUNK, p.kv_len - chunk_start);
|
||||||
|
|
||||||
|
int total = this_chunk * p.head_dim;
|
||||||
|
for (int i = threadIdx.y * 32 + lane; i < total; i += blockDim.x * blockDim.y)
|
||||||
|
k_smem[i] = p.k[kv_base + chunk_start * p.head_dim + i];
|
||||||
|
__syncthreads();
|
||||||
|
|
||||||
|
for (int s = 0; s < this_chunk; s++) {
|
||||||
|
float partial = 0.0f;
|
||||||
|
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 = warp_reduce_sum(partial) * p.scale;
|
||||||
|
|
||||||
|
if (p.use_mask && p.mask && !p.mask[mask_base + chunk_start + s])
|
||||||
|
partial = -FLT_MAX;
|
||||||
|
if (p.is_causal && (chunk_start + s) > p.causal_offset)
|
||||||
|
partial = -FLT_MAX;
|
||||||
|
|
||||||
|
float new_m = fmaxf(m, partial);
|
||||||
|
float alpha = expf(m - new_m);
|
||||||
|
float beta = expf(partial - new_m);
|
||||||
|
d = d * alpha + beta;
|
||||||
|
|
||||||
|
int v_off = kv_base + (chunk_start + s) * p.head_dim + lane * hd_per_thread;
|
||||||
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
|
acc_reg[i] = acc_reg[i] * alpha + __bfloat162float(p.v[v_off + i]) * beta;
|
||||||
|
m = new_m;
|
||||||
|
}
|
||||||
|
__syncthreads();
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- write UN-normalised partials for this split ----
|
||||||
|
size_t bh = (size_t)batch * p.q_head + q_head;
|
||||||
|
size_t slot = bh * p.num_splits + split;
|
||||||
|
int d0 = lane * hd_per_thread;
|
||||||
|
for (int i = 0; i < hd_per_thread; i++) {
|
||||||
|
int dd = d0 + i;
|
||||||
|
p.o_part[slot * p.head_dim + dd] = acc_reg[i];
|
||||||
|
}
|
||||||
|
if (lane == 0) {
|
||||||
|
p.ml_part[slot * 2] = m;
|
||||||
|
p.ml_part[slot * 2 + 1] = d;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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) {
|
||||||
|
int bh = blockIdx.x;
|
||||||
|
int d = threadIdx.x;
|
||||||
|
if (d >= p.head_dim) return;
|
||||||
|
|
||||||
|
size_t split_base = (size_t)bh * p.num_splits;
|
||||||
|
const float* mlp = p.ml_part + split_base * 2;
|
||||||
|
const float* op = p.o_part + split_base * p.head_dim;
|
||||||
|
|
||||||
|
float m = -FLT_MAX, l = 0.0f, acc = 0.0f;
|
||||||
|
for (int s = 0; s < p.num_splits; s++) {
|
||||||
|
float mi = mlp[s * 2];
|
||||||
|
if (mi <= -FLT_MAX) continue;
|
||||||
|
float li = mlp[s * 2 + 1];
|
||||||
|
float nm = fmaxf(m, mi);
|
||||||
|
float corr = __expf(m - nm);
|
||||||
|
float e = __expf(mi - nm);
|
||||||
|
acc = acc * corr + op[s * p.head_dim + d] * e;
|
||||||
|
l = l * corr + li * e;
|
||||||
|
m = nm;
|
||||||
|
}
|
||||||
|
|
||||||
|
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
||||||
|
p.o[(size_t)bh * p.head_dim + d] = __float2bfloat16(acc * inv);
|
||||||
|
}
|
||||||
@@ -0,0 +1,200 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cfloat>
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
#include "attn_mma_utils.cuh"
|
||||||
|
|
||||||
|
using bf16 = __nv_bfloat16;
|
||||||
|
|
||||||
|
// Split-K (FlashDecoding) tensor-core decode via GQA head-packing.
|
||||||
|
//
|
||||||
|
// Decode has q_len == 1, so S = q @ K^T is a GEMV per head — no tensor-core work
|
||||||
|
// on its own. But GQA gives us G = q_head / kv_head query heads that all 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 reuses each loaded K/V tile
|
||||||
|
// across all G heads (K/V load is the decode 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.
|
||||||
|
//
|
||||||
|
// Partial layout (float, contiguous):
|
||||||
|
// o_part : [batch, q_head, num_splits, HEAD_DIM]
|
||||||
|
// ml_part: [batch, q_head, num_splits, 2] (m, l)
|
||||||
|
//
|
||||||
|
// Optimizations:
|
||||||
|
// - cp.async global→shared for K/V (bypasses registers, cuts instruction count)
|
||||||
|
// - XOR swizzle (swiz_col): LD=HEAD_DIM, zero waste, no bank conflicts
|
||||||
|
// - Q loaded directly from global into mma A-operand registers (no sQ staging,
|
||||||
|
// no prologue syncwarp) — frees shared memory for double-buffering
|
||||||
|
// - Double-buffered KV (STAGES=2): next tile's cp.async overlaps current
|
||||||
|
// tile's MMA compute — hides global load latency / boosts bandwidth
|
||||||
|
// utilization for small-batch (low-occupancy) decode
|
||||||
|
// - Predicated cp.async (cp_async_16_pred) for full AND partial tiles on one
|
||||||
|
// uniform path — eliminates the scalar fallback branch
|
||||||
|
//
|
||||||
|
// Smem footprint (BC=32): STAGES=2 → 2*(sK+sV) = 2*2*32*HEAD_DIM*2 bytes.
|
||||||
|
// D=128: 16 KB (fits 48 KB static cap). D=256: 32 KB (also fits).
|
||||||
|
// STAGES=1 fallback (4/8 KB) for smem-constrained configs.
|
||||||
|
template <int HEAD_DIM, int BC, int STAGES = 2>
|
||||||
|
__global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
||||||
|
constexpr int BR = 16;
|
||||||
|
constexpr int KD = HEAD_DIM / 16;
|
||||||
|
constexpr int NC8 = BC / 8;
|
||||||
|
constexpr int KT2 = BC / 16;
|
||||||
|
constexpr int DN8 = HEAD_DIM / 8;
|
||||||
|
constexpr int LD = HEAD_DIM;
|
||||||
|
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
|
||||||
|
constexpr int VEC = 8;
|
||||||
|
constexpr int TOTAL = BC * HEAD_DIM;
|
||||||
|
|
||||||
|
const int lane = threadIdx.x;
|
||||||
|
const int gid = lane >> 2;
|
||||||
|
const int tid4 = lane & 3;
|
||||||
|
|
||||||
|
const int kv_head = blockIdx.x;
|
||||||
|
const int batch = blockIdx.y;
|
||||||
|
const int split = blockIdx.z;
|
||||||
|
const int G = p.q_head / p.kv_head;
|
||||||
|
const int q_head0 = kv_head * G;
|
||||||
|
|
||||||
|
// Double-buffered shared memory for K/V (no sQ needed — Q goes direct
|
||||||
|
// from global to registers).
|
||||||
|
__shared__ __align__(16) bf16 sK[STAGES * BC * LD];
|
||||||
|
__shared__ __align__(16) bf16 sV[STAGES * BC * LD];
|
||||||
|
|
||||||
|
// ---- Load Q directly from global into mma A-operand registers ----
|
||||||
|
// Same layout as prefill: frag[0]/[2] = row gid, frag[1]/[3] = row gid+8
|
||||||
|
// cols kt*16 + tid4*2 + {0,1} / +{8,9}. pau[0]=cols c,c+1; pau[4]=c+8,c+9.
|
||||||
|
const int q_base = (batch * p.q_head + q_head0) * HEAD_DIM;
|
||||||
|
const int qra = gid;
|
||||||
|
const int qrb = gid + 8;
|
||||||
|
const bool va = qra < G, vb = qrb < G;
|
||||||
|
unsigned Qa[KD][4];
|
||||||
|
#pragma unroll
|
||||||
|
for (int kt = 0; kt < KD; kt++) {
|
||||||
|
int c = kt * 16 + tid4 * 2;
|
||||||
|
const unsigned* pau = reinterpret_cast<const unsigned*>(
|
||||||
|
&p.q[q_base + qra * HEAD_DIM + c]);
|
||||||
|
const unsigned* pbu = reinterpret_cast<const unsigned*>(
|
||||||
|
&p.q[q_base + qrb * HEAD_DIM + c]);
|
||||||
|
Qa[kt][0] = va ? pau[0] : 0u;
|
||||||
|
Qa[kt][1] = vb ? pbu[0] : 0u;
|
||||||
|
Qa[kt][2] = va ? pau[4] : 0u;
|
||||||
|
Qa[kt][3] = vb ? pbu[4] : 0u;
|
||||||
|
}
|
||||||
|
|
||||||
|
float Oacc[DN8][4];
|
||||||
|
#pragma unroll
|
||||||
|
for (int j = 0; j < DN8; j++)
|
||||||
|
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
||||||
|
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
||||||
|
|
||||||
|
const int kv_base = (batch * p.kv_head + kv_head) * p.kv_len * HEAD_DIM;
|
||||||
|
const int mask_base = batch * p.kv_len;
|
||||||
|
const int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||||
|
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
|
||||||
|
const int ti_begin = split * tiles_per_split;
|
||||||
|
const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
|
||||||
|
const int has_mask = p.use_mask && p.mask;
|
||||||
|
|
||||||
|
// ---- Load tile lambda: predicated cp.async, unified full/partial ----
|
||||||
|
auto load_tile = [&](int ti, int buf) {
|
||||||
|
int kv0 = ti * BC;
|
||||||
|
bf16* dK = sK + buf * BC * LD;
|
||||||
|
bf16* dV = sV + buf * BC * LD;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = lane * VEC; i < TOTAL; i += 32 * VEC) {
|
||||||
|
int r = i / HEAD_DIM, d = i % HEAD_DIM;
|
||||||
|
int kc = kv0 + r;
|
||||||
|
bool valid = kc < p.kv_len;
|
||||||
|
int off = r * LD + swiz_col(d, r, SWIZ_MASK);
|
||||||
|
cp_async_16_pred(&dK[off], &p.k[kv_base + kc * HEAD_DIM + d], valid);
|
||||||
|
cp_async_16_pred(&dV[off], &p.v[kv_base + kc * HEAD_DIM + d], valid);
|
||||||
|
}
|
||||||
|
cp_async_commit();
|
||||||
|
};
|
||||||
|
|
||||||
|
// ---- Prologue: issue first tile load ----
|
||||||
|
if (ti_begin < ti_end) {
|
||||||
|
load_tile(ti_begin, 0);
|
||||||
|
}
|
||||||
|
|
||||||
|
for (int ti = ti_begin; ti < ti_end; ti++) {
|
||||||
|
constexpr int BUF_MASK = (STAGES > 1) ? (STAGES - 1) : 0;
|
||||||
|
int buf = (ti - ti_begin) & BUF_MASK;
|
||||||
|
|
||||||
|
// 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>();
|
||||||
|
__syncwarp();
|
||||||
|
if constexpr (STAGES > 1) {
|
||||||
|
if (ti + 1 < ti_end)
|
||||||
|
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
|
||||||
|
}
|
||||||
|
|
||||||
|
const bf16* bK = sK + buf * BC * LD;
|
||||||
|
const bf16* bV = sV + buf * BC * LD;
|
||||||
|
int kv0 = ti * BC;
|
||||||
|
|
||||||
|
float Sacc[NC8][4];
|
||||||
|
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
|
||||||
|
|
||||||
|
#pragma unroll
|
||||||
|
for (int n8 = 0; n8 < NC8; n8++)
|
||||||
|
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||||
|
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||||
|
|
||||||
|
int maxc = p.is_causal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
||||||
|
mma_softmax_tile<NC8, DN8>(kv0, maxc, maxc,
|
||||||
|
mask_base, p.mask, has_mask,
|
||||||
|
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||||
|
|
||||||
|
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc);
|
||||||
|
__syncwarp();
|
||||||
|
|
||||||
|
if constexpr (STAGES == 1) {
|
||||||
|
if (ti + 1 < ti_end)
|
||||||
|
load_tile(ti + 1, 0);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- write UN-normalised partials for this split ----
|
||||||
|
auto split_slot = [&](int h) -> size_t {
|
||||||
|
size_t bh = (size_t)batch * p.q_head + h;
|
||||||
|
return bh * p.num_splits + split;
|
||||||
|
};
|
||||||
|
#pragma unroll
|
||||||
|
for (int dn8 = 0; dn8 < DN8; dn8++) {
|
||||||
|
int d = dn8 * 8 + 2 * tid4;
|
||||||
|
int r0 = gid, r1 = gid + 8;
|
||||||
|
if (r0 < G) {
|
||||||
|
int h = q_head0 + r0;
|
||||||
|
float* op = p.o_part + split_slot(h) * HEAD_DIM;
|
||||||
|
op[d] = Oacc[dn8][0];
|
||||||
|
op[d + 1] = Oacc[dn8][1];
|
||||||
|
}
|
||||||
|
if (r1 < G) {
|
||||||
|
int h = q_head0 + r1;
|
||||||
|
float* op = p.o_part + split_slot(h) * HEAD_DIM;
|
||||||
|
op[d] = Oacc[dn8][2];
|
||||||
|
op[d + 1] = Oacc[dn8][3];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (tid4 == 0) {
|
||||||
|
int r0 = gid, r1 = gid + 8;
|
||||||
|
if (r0 < G) {
|
||||||
|
int h = q_head0 + r0;
|
||||||
|
float* mp = p.ml_part + split_slot(h) * 2;
|
||||||
|
mp[0] = m0; mp[1] = l0;
|
||||||
|
}
|
||||||
|
if (r1 < G) {
|
||||||
|
int h = q_head0 + r1;
|
||||||
|
float* mp = p.ml_part + split_slot(h) * 2;
|
||||||
|
mp[0] = m1; mp[1] = l1;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,43 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <torch/extension.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
|
||||||
|
template<typename T>
|
||||||
|
inline void attn_pack_params(
|
||||||
|
torch::Tensor q,
|
||||||
|
torch::Tensor k,
|
||||||
|
torch::Tensor v,
|
||||||
|
c10::optional<torch::Tensor> mask,
|
||||||
|
bool is_causal,
|
||||||
|
int64_t causal_offset,
|
||||||
|
c10::optional<double> scale,
|
||||||
|
AttentionParams<T>& p
|
||||||
|
) {
|
||||||
|
TORCH_CHECK(q.is_cuda() && k.is_cuda() && v.is_cuda());
|
||||||
|
TORCH_CHECK(q.dtype() == torch::kBFloat16);
|
||||||
|
TORCH_CHECK(k.dtype() == torch::kBFloat16);
|
||||||
|
TORCH_CHECK(v.dtype() == torch::kBFloat16);
|
||||||
|
|
||||||
|
p.batch = (int)q.size(0);
|
||||||
|
p.q_head = (int)q.size(1);
|
||||||
|
p.kv_head = (int)k.size(1);
|
||||||
|
p.q_len = (int)q.size(2);
|
||||||
|
p.kv_len = (int)k.size(2);
|
||||||
|
p.head_dim = (int)q.size(3);
|
||||||
|
p.use_mask = mask.has_value() ? 1 : 0;
|
||||||
|
p.is_causal = is_causal ? 1 : 0;
|
||||||
|
p.causal_offset = (int)causal_offset;
|
||||||
|
p.scale = scale.has_value() ? (float)scale.value() : 1.0f / sqrtf((float)p.head_dim);
|
||||||
|
p.q = (const T*)q.data_ptr();
|
||||||
|
p.k = (const T*)k.data_ptr();
|
||||||
|
p.v = (const T*)v.data_ptr();
|
||||||
|
if (p.use_mask) {
|
||||||
|
TORCH_CHECK(mask.value().dtype() == torch::kBool);
|
||||||
|
TORCH_CHECK(mask.value().dim() == 2);
|
||||||
|
TORCH_CHECK(mask.value().size(0) == p.batch);
|
||||||
|
TORCH_CHECK(mask.value().size(1) == p.kv_len);
|
||||||
|
p.mask = mask.value().data_ptr<bool>();
|
||||||
|
} else {
|
||||||
|
p.mask = nullptr;
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,252 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cfloat>
|
||||||
|
#include <cuda_fp16.h>
|
||||||
|
#include <cuda_runtime.h>
|
||||||
|
|
||||||
|
// Shared MMA utilities for tensor-core GQA kernels.
|
||||||
|
// mma.sync.m16n8k16 PTX wrappers, ldmatrix helpers, and bf16 packing.
|
||||||
|
|
||||||
|
// mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32
|
||||||
|
__device__ __forceinline__ void mma16816(float* d, const unsigned* a,
|
||||||
|
const unsigned* b, const float* c) {
|
||||||
|
asm volatile(
|
||||||
|
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
|
||||||
|
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};"
|
||||||
|
: "=f"(d[0]), "=f"(d[1]), "=f"(d[2]), "=f"(d[3])
|
||||||
|
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]),
|
||||||
|
"f"(c[0]), "f"(c[1]), "f"(c[2]), "f"(c[3]));
|
||||||
|
}
|
||||||
|
|
||||||
|
// read two adjacent bf16 from smem as one packed .b32 (elem0 low, elem1 high)
|
||||||
|
__device__ __forceinline__ unsigned ld2(const bf16* p) {
|
||||||
|
return *reinterpret_cast<const unsigned*>(p);
|
||||||
|
}
|
||||||
|
|
||||||
|
// pack two floats into one bf16x2 as .b32
|
||||||
|
__device__ __forceinline__ unsigned pk2(float a, float b) {
|
||||||
|
__nv_bfloat162 v = __floats2bfloat162_rn(a, b);
|
||||||
|
return *reinterpret_cast<unsigned*>(&v);
|
||||||
|
}
|
||||||
|
|
||||||
|
// pack two (non-contiguous) bf16 into one .b32
|
||||||
|
__device__ __forceinline__ unsigned pkb(bf16 a, bf16 b) {
|
||||||
|
__nv_bfloat162 v;
|
||||||
|
v.x = a;
|
||||||
|
v.y = b;
|
||||||
|
return *reinterpret_cast<unsigned*>(&v);
|
||||||
|
}
|
||||||
|
|
||||||
|
// ldmatrix: cooperatively load mma fragments from smem (one instruction per
|
||||||
|
// 16x16 / 16x8 tile) with the exact register layout mma expects — replaces the
|
||||||
|
// 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) {
|
||||||
|
unsigned a = __cvta_generic_to_shared(p);
|
||||||
|
asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];"
|
||||||
|
: "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3])
|
||||||
|
: "r"(a));
|
||||||
|
}
|
||||||
|
__device__ __forceinline__ void ldmatrix_x2(unsigned* r, const bf16* p) {
|
||||||
|
unsigned a = __cvta_generic_to_shared(p);
|
||||||
|
asm volatile("ldmatrix.sync.aligned.m8n8.x2.shared.b16 {%0,%1}, [%2];"
|
||||||
|
: "=r"(r[0]), "=r"(r[1])
|
||||||
|
: "r"(a));
|
||||||
|
}
|
||||||
|
__device__ __forceinline__ void ldmatrix_x2_trans(unsigned* r, const bf16* p) {
|
||||||
|
unsigned a = __cvta_generic_to_shared(p);
|
||||||
|
asm volatile("ldmatrix.sync.aligned.m8n8.x2.trans.shared.b16 {%0,%1}, [%2];"
|
||||||
|
: "=r"(r[0]), "=r"(r[1])
|
||||||
|
: "r"(a));
|
||||||
|
}
|
||||||
|
|
||||||
|
// XOR swizzle for shared-memory column at 8-bf16 chunk granularity.
|
||||||
|
// 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) {
|
||||||
|
return ((d >> 3) ^ (r & mask)) << 3 | (d & 7);
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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) {
|
||||||
|
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
|
||||||
|
asm volatile("cp.async.ca.shared.global [%0], [%1], 16;"
|
||||||
|
:: "r"(smem_addr), "l"(gmem_ptr));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill the
|
||||||
|
// destination (src-size operand = 0 → no bytes read from src, so an
|
||||||
|
// 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,
|
||||||
|
const void* gmem_ptr,
|
||||||
|
bool pred) {
|
||||||
|
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
|
||||||
|
int src_size = pred ? 16 : 0;
|
||||||
|
asm volatile("cp.async.ca.shared.global [%0], [%1], 16, %2;"
|
||||||
|
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
|
||||||
|
}
|
||||||
|
|
||||||
|
__device__ __forceinline__ void cp_async_commit() {
|
||||||
|
asm volatile("cp.async.commit_group;");
|
||||||
|
}
|
||||||
|
|
||||||
|
__device__ __forceinline__ void cp_async_wait_all() {
|
||||||
|
asm volatile("cp.async.wait_all;");
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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>
|
||||||
|
__device__ __forceinline__ void cp_async_wait_group() {
|
||||||
|
asm volatile("cp.async.wait_group %0;" :: "n"(N));
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
// 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
|
||||||
|
// caller to avoid bf16 precision loss).
|
||||||
|
// LD and SWIZ_MASK are constexpr in the calling kernel — passing them as
|
||||||
|
// runtime ints lets the compiler fold them while keeping the signature clean.
|
||||||
|
template <int KD, int NC8>
|
||||||
|
__device__ inline void mma_compute_scores(
|
||||||
|
const unsigned Qa[KD][4],
|
||||||
|
const bf16* __restrict__ sK,
|
||||||
|
int LD, int SWIZ_MASK, int lane,
|
||||||
|
float Sacc[NC8][4])
|
||||||
|
{
|
||||||
|
#pragma unroll
|
||||||
|
for (int n8 = 0; n8 < NC8; n8++) {
|
||||||
|
Sacc[n8][0] = Sacc[n8][1] = Sacc[n8][2] = Sacc[n8][3] = 0.0f;
|
||||||
|
int krow_l = n8 * 8 + (lane & 7);
|
||||||
|
int kcol_h = (lane & 8) ? 8 : 0;
|
||||||
|
#pragma unroll
|
||||||
|
for (int kt = 0; kt < KD; kt++) {
|
||||||
|
unsigned b[2];
|
||||||
|
ldmatrix_x2(b, &sK[krow_l * LD + swiz_col(kt * 16 + kcol_h, krow_l, SWIZ_MASK)]);
|
||||||
|
mma16816(Sacc[n8], Qa[kt], b, Sacc[n8]);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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).
|
||||||
|
// Reads Sacc (Q@K^T scores), applies causal/mask, computes P = exp(S - nm),
|
||||||
|
// 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(
|
||||||
|
int kv0,
|
||||||
|
int maxc0, int maxc1,
|
||||||
|
int mask_base,
|
||||||
|
const bool* __restrict__ mask,
|
||||||
|
bool has_mask,
|
||||||
|
float Sacc[NC8][4],
|
||||||
|
float Oacc[DN8][4],
|
||||||
|
float& m0, float& m1,
|
||||||
|
float& l0, float& l1,
|
||||||
|
int lane)
|
||||||
|
{
|
||||||
|
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;
|
||||||
|
#pragma unroll
|
||||||
|
for (int n8 = 0; n8 < NC8; n8++) {
|
||||||
|
int cc = kv0 + n8 * 8 + 2 * tid4;
|
||||||
|
int c1 = cc + 1;
|
||||||
|
bool b0 = (cc >= maxc0) || (has_mask && !mask[mask_base + cc]);
|
||||||
|
bool b1 = (c1 >= maxc0) || (has_mask && !mask[mask_base + c1]);
|
||||||
|
bool b2 = (cc >= maxc1) || (has_mask && !mask[mask_base + cc]);
|
||||||
|
bool b3 = (c1 >= maxc1) || (has_mask && !mask[mask_base + c1]);
|
||||||
|
float s0 = b0 ? -FLT_MAX : Sacc[n8][0];
|
||||||
|
float s1 = b1 ? -FLT_MAX : Sacc[n8][1];
|
||||||
|
float s2 = b2 ? -FLT_MAX : Sacc[n8][2];
|
||||||
|
float s3 = b3 ? -FLT_MAX : Sacc[n8][3];
|
||||||
|
Sacc[n8][0] = s0; Sacc[n8][1] = s1;
|
||||||
|
Sacc[n8][2] = s2; Sacc[n8][3] = s3;
|
||||||
|
rmax0 = fmaxf(rmax0, fmaxf(s0, s1));
|
||||||
|
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, 2));
|
||||||
|
rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 1));
|
||||||
|
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);
|
||||||
|
// 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 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 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;
|
||||||
|
#pragma unroll
|
||||||
|
for (int n8 = 0; n8 < NC8; n8++) {
|
||||||
|
float p0 = pn0 * __expf(Sacc[n8][0] - nm0);
|
||||||
|
float p1 = pn0 * __expf(Sacc[n8][1] - nm0);
|
||||||
|
float p2 = pn1 * __expf(Sacc[n8][2] - nm1);
|
||||||
|
float p3 = pn1 * __expf(Sacc[n8][3] - nm1);
|
||||||
|
Sacc[n8][0] = p0; Sacc[n8][1] = p1;
|
||||||
|
Sacc[n8][2] = p2; Sacc[n8][3] = p3;
|
||||||
|
rsum0 += p0 + p1;
|
||||||
|
rsum1 += p2 + p3;
|
||||||
|
}
|
||||||
|
rsum0 += __shfl_xor_sync(0xFFFFFFFF, rsum0, 1);
|
||||||
|
rsum0 += __shfl_xor_sync(0xFFFFFFFF, rsum0, 2);
|
||||||
|
rsum1 += __shfl_xor_sync(0xFFFFFFFF, rsum1, 1);
|
||||||
|
rsum1 += __shfl_xor_sync(0xFFFFFFFF, rsum1, 2);
|
||||||
|
l0 = l0 * corr0 + rsum0;
|
||||||
|
l1 = l1 * corr1 + rsum1;
|
||||||
|
m0 = nm0; m1 = nm1;
|
||||||
|
|
||||||
|
#pragma unroll
|
||||||
|
for (int j = 0; j < DN8; j++) {
|
||||||
|
Oacc[j][0] *= corr0; Oacc[j][1] *= corr0;
|
||||||
|
Oacc[j][2] *= corr1; Oacc[j][3] *= corr1;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// O += P @ V (Sacc must contain P = attention weights after softmax).
|
||||||
|
template <int DN8, int KT2>
|
||||||
|
__device__ inline void mma_pv_accumulate(
|
||||||
|
float Sacc[][4],
|
||||||
|
const bf16* __restrict__ sV,
|
||||||
|
int LD, int SWIZ_MASK, int lane,
|
||||||
|
float Oacc[DN8][4])
|
||||||
|
{
|
||||||
|
#pragma unroll
|
||||||
|
for (int kt2 = 0; kt2 < KT2; kt2++) {
|
||||||
|
unsigned Pa[4];
|
||||||
|
Pa[0] = pk2(Sacc[kt2 * 2][0], Sacc[kt2 * 2][1]);
|
||||||
|
Pa[1] = pk2(Sacc[kt2 * 2][2], Sacc[kt2 * 2][3]);
|
||||||
|
Pa[2] = pk2(Sacc[kt2 * 2 + 1][0], Sacc[kt2 * 2 + 1][1]);
|
||||||
|
Pa[3] = pk2(Sacc[kt2 * 2 + 1][2], Sacc[kt2 * 2 + 1][3]);
|
||||||
|
int vrow_l = kt2 * 16 + (lane & 15);
|
||||||
|
#pragma unroll
|
||||||
|
for (int dn8 = 0; dn8 < DN8; dn8++) {
|
||||||
|
unsigned b[2];
|
||||||
|
ldmatrix_x2_trans(b, &sV[vrow_l * LD + swiz_col(dn8 * 8, vrow_l, SWIZ_MASK)]);
|
||||||
|
mma16816(Oacc[dn8], Pa, b, Oacc[dn8]);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,149 @@
|
|||||||
|
#include "attn_paged_decode_split_kv.cuh"
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
#include "attn_paged_decode_split_kv_mma.cuh"
|
||||||
|
#endif
|
||||||
|
|
||||||
|
#include <torch/extension.h>
|
||||||
|
#include <c10/cuda/CUDAGuard.h>
|
||||||
|
|
||||||
|
static int paged_decode_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)));
|
||||||
|
}
|
||||||
|
|
||||||
|
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 = paged_decode_num_splits(p.batch * p.kv_head, chunks_total);
|
||||||
|
|
||||||
|
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 ml_part = torch::empty({p.batch, p.q_head, p.num_splits, 2}, fopt);
|
||||||
|
p.o_part = o_part.data_ptr<float>();
|
||||||
|
p.ml_part = ml_part.data_ptr<float>();
|
||||||
|
|
||||||
|
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 = paged_decode_num_splits(p.batch * p.kv_head, tiles_total);
|
||||||
|
|
||||||
|
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 ml_part = torch::empty({p.batch, p.q_head, p.num_splits, 2}, fopt);
|
||||||
|
p.o_part = o_part.data_ptr<float>();
|
||||||
|
p.ml_part = ml_part.data_ptr<float>();
|
||||||
|
|
||||||
|
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 (!p.use_mask && 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 q,
|
||||||
|
torch::Tensor page_table,
|
||||||
|
torch::Tensor k_cache,
|
||||||
|
torch::Tensor v_cache,
|
||||||
|
int64_t page_size,
|
||||||
|
int64_t kv_len,
|
||||||
|
c10::optional<torch::Tensor> mask,
|
||||||
|
bool is_causal = false,
|
||||||
|
int64_t causal_offset = 0,
|
||||||
|
c10::optional<double> scale = c10::nullopt
|
||||||
|
) {
|
||||||
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
|
|
||||||
|
int batch = q.size(0);
|
||||||
|
int q_head = q.size(1);
|
||||||
|
int head_dim = q.size(3);
|
||||||
|
int kv_head = k_cache.size(2);
|
||||||
|
int max_pages = page_table.size(1);
|
||||||
|
|
||||||
|
TORCH_CHECK(q.is_cuda() && page_table.is_cuda() && k_cache.is_cuda() && v_cache.is_cuda());
|
||||||
|
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
|
||||||
|
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16");
|
||||||
|
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16");
|
||||||
|
TORCH_CHECK(page_table.dtype() == torch::kLong, "page_table must be int64");
|
||||||
|
TORCH_CHECK(q.size(2) == 1, "Q seq_len must be 1 (decode)");
|
||||||
|
TORCH_CHECK(head_dim % 32 == 0, "head_dim must be multiple of 32");
|
||||||
|
TORCH_CHECK(k_cache.size(1) == page_size,
|
||||||
|
"k_cache dim 1 must equal page_size, got ",
|
||||||
|
k_cache.size(1), " vs ", page_size);
|
||||||
|
TORCH_CHECK(k_cache.size(0) >= 0, "k_cache must have at least 0 pages");
|
||||||
|
|
||||||
|
float scale_val = scale.has_value()
|
||||||
|
? static_cast<float>(scale.value())
|
||||||
|
: 1.0f / std::sqrt(static_cast<float>(head_dim));
|
||||||
|
|
||||||
|
auto O = torch::empty_like(q);
|
||||||
|
|
||||||
|
PagedAttentionParams<bf16, float> p;
|
||||||
|
p.batch = batch;
|
||||||
|
p.q_head = q_head;
|
||||||
|
p.kv_head = kv_head;
|
||||||
|
p.q_len = static_cast<int>(q.size(2));
|
||||||
|
p.kv_len = static_cast<int>(kv_len);
|
||||||
|
p.head_dim = head_dim;
|
||||||
|
p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0;
|
||||||
|
p.is_causal = is_causal ? 1 : 0;
|
||||||
|
p.causal_offset = static_cast<int>(causal_offset);
|
||||||
|
p.scale = scale_val;
|
||||||
|
p.page_size = static_cast<int>(page_size);
|
||||||
|
p.max_pages = max_pages;
|
||||||
|
p.page_table = page_table.data_ptr<int64_t>();
|
||||||
|
p.k_cache = reinterpret_cast<const bf16*>(k_cache.data_ptr());
|
||||||
|
p.v_cache = reinterpret_cast<const bf16*>(v_cache.data_ptr());
|
||||||
|
p.q = reinterpret_cast<const bf16*>(q.data_ptr());
|
||||||
|
p.mask = p.use_mask ? mask.value().data_ptr<bool>() : nullptr;
|
||||||
|
p.o = reinterpret_cast<bf16*>(O.data_ptr());
|
||||||
|
p.o_part = nullptr;
|
||||||
|
p.ml_part = nullptr;
|
||||||
|
|
||||||
|
switch (p.head_dim) {
|
||||||
|
case 32: dispatch_paged_decode<32>(p); break;
|
||||||
|
case 64: dispatch_paged_decode<64>(p); break;
|
||||||
|
case 128: dispatch_paged_decode<128>(p); break;
|
||||||
|
case 256: dispatch_paged_decode<256>(p); break;
|
||||||
|
default:
|
||||||
|
TORCH_CHECK(false, "paged_decode: unsupported head_dim ", p.head_dim,
|
||||||
|
" (supported: 32, 64, 128, 256)");
|
||||||
|
}
|
||||||
|
|
||||||
|
return O;
|
||||||
|
}
|
||||||
|
|
||||||
|
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||||
|
m.def("attn_paged_decode", &attn_paged_decode,
|
||||||
|
py::arg("q"),
|
||||||
|
py::arg("page_table"),
|
||||||
|
py::arg("k_cache"),
|
||||||
|
py::arg("v_cache"),
|
||||||
|
py::arg("page_size"),
|
||||||
|
py::arg("kv_len"),
|
||||||
|
py::arg("mask") = py::none(),
|
||||||
|
py::arg("is_causal") = false,
|
||||||
|
py::arg("causal_offset") = 0,
|
||||||
|
py::arg("scale") = py::none(),
|
||||||
|
"Paged GQA decode — split-KV with direct page-table access.");
|
||||||
|
}
|
||||||
@@ -0,0 +1,140 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include <float.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
|
||||||
|
using bf16 = __nv_bfloat16;
|
||||||
|
constexpr int PDC_CHUNK = 64;
|
||||||
|
|
||||||
|
__device__ inline float paged_warp_reduce_sum(float val) {
|
||||||
|
for (int offset = 16; offset > 0; offset >>= 1)
|
||||||
|
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
|
||||||
|
return val;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Split-KV scalar decode: one warp per query head, grid.z partitions KV.
|
||||||
|
__global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p) {
|
||||||
|
int batch = blockIdx.x / p.kv_head;
|
||||||
|
int kv_head = blockIdx.x % p.kv_head;
|
||||||
|
int split = blockIdx.z;
|
||||||
|
int group_size = blockDim.y;
|
||||||
|
int q_head = kv_head * group_size + threadIdx.y;
|
||||||
|
int lane = threadIdx.x;
|
||||||
|
int hd_per_thread = p.head_dim / 32;
|
||||||
|
|
||||||
|
float q_reg[8];
|
||||||
|
int q_off = ((batch * p.q_head + q_head) * 1) * p.head_dim + lane * hd_per_thread;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
|
q_reg[i] = __bfloat162float(p.q[q_off + i]);
|
||||||
|
|
||||||
|
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
|
||||||
|
|
||||||
|
extern __shared__ __align__(16) bf16 k_smem[];
|
||||||
|
|
||||||
|
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
||||||
|
int chunks_per_split = (chunks_total + p.num_splits - 1) / p.num_splits;
|
||||||
|
int ch_begin = split * chunks_per_split;
|
||||||
|
int ch_end = min(chunks_total, ch_begin + chunks_per_split);
|
||||||
|
|
||||||
|
const int mask_base = batch * p.kv_len;
|
||||||
|
|
||||||
|
for (int ci = ch_begin; ci < ch_end; ci++) {
|
||||||
|
int chunk_start = ci * PDC_CHUNK;
|
||||||
|
int this_chunk = min(PDC_CHUNK, p.kv_len - chunk_start);
|
||||||
|
|
||||||
|
int total = this_chunk * p.head_dim;
|
||||||
|
for (int i = threadIdx.y * 32 + lane; i < total; i += blockDim.x * blockDim.y) {
|
||||||
|
int s = i / p.head_dim;
|
||||||
|
int d_dim = i % p.head_dim;
|
||||||
|
int pos = chunk_start + s;
|
||||||
|
int logical_page = pos / p.page_size;
|
||||||
|
int page_offset = pos % p.page_size;
|
||||||
|
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
||||||
|
if (phys_page >= 0) {
|
||||||
|
int64_t off = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim
|
||||||
|
+ (int64_t)page_offset * p.kv_head * p.head_dim
|
||||||
|
+ (int64_t)kv_head * p.head_dim
|
||||||
|
+ d_dim;
|
||||||
|
k_smem[i] = p.k_cache[off];
|
||||||
|
} else {
|
||||||
|
k_smem[i] = __float2bfloat16(0.0f);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
__syncthreads();
|
||||||
|
|
||||||
|
for (int s = 0; s < this_chunk; s++) {
|
||||||
|
float partial = 0.0f;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
|
partial += q_reg[i] * __bfloat162float(k_smem[s * p.head_dim + lane * hd_per_thread + i]);
|
||||||
|
partial = paged_warp_reduce_sum(partial) * p.scale;
|
||||||
|
|
||||||
|
if (p.use_mask && p.mask && !p.mask[mask_base + chunk_start + s])
|
||||||
|
partial = -FLT_MAX;
|
||||||
|
if (p.is_causal && (chunk_start + s) > p.causal_offset)
|
||||||
|
partial = -FLT_MAX;
|
||||||
|
|
||||||
|
float new_m = fmaxf(m, partial);
|
||||||
|
float alpha = expf(m - new_m);
|
||||||
|
float beta = expf(partial - new_m);
|
||||||
|
d = d * alpha + beta;
|
||||||
|
|
||||||
|
int pos = chunk_start + s;
|
||||||
|
int logical_page = pos / p.page_size;
|
||||||
|
int page_offset = pos % p.page_size;
|
||||||
|
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
||||||
|
if (phys_page >= 0) {
|
||||||
|
int64_t v_base = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim
|
||||||
|
+ (int64_t)page_offset * p.kv_head * p.head_dim
|
||||||
|
+ (int64_t)kv_head * p.head_dim;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
|
acc_reg[i] = acc_reg[i] * alpha + __bfloat162float(p.v_cache[v_base + lane * hd_per_thread + i]) * beta;
|
||||||
|
} else {
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
|
acc_reg[i] = acc_reg[i] * alpha + 0.0f * beta;
|
||||||
|
}
|
||||||
|
m = new_m;
|
||||||
|
}
|
||||||
|
__syncthreads();
|
||||||
|
}
|
||||||
|
|
||||||
|
size_t bh = (size_t)batch * p.q_head + q_head;
|
||||||
|
size_t slot = bh * p.num_splits + split;
|
||||||
|
int d0 = lane * hd_per_thread;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
|
p.o_part[slot * p.head_dim + (d0 + i)] = acc_reg[i];
|
||||||
|
if (lane == 0) {
|
||||||
|
p.ml_part[slot * 2] = m;
|
||||||
|
p.ml_part[slot * 2 + 1] = d;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
__global__ void paged_attn_decode_combine_kernel(PagedAttentionParams<bf16> p) {
|
||||||
|
int bh = blockIdx.x;
|
||||||
|
int d = threadIdx.x;
|
||||||
|
if (d >= p.head_dim) return;
|
||||||
|
|
||||||
|
size_t split_base = (size_t)bh * p.num_splits;
|
||||||
|
const float* mlp = p.ml_part + split_base * 2;
|
||||||
|
const float* op = p.o_part + split_base * p.head_dim;
|
||||||
|
|
||||||
|
float m = -FLT_MAX, l = 0.0f, acc = 0.0f;
|
||||||
|
for (int s = 0; s < p.num_splits; s++) {
|
||||||
|
float mi = mlp[s * 2];
|
||||||
|
if (mi <= -FLT_MAX) continue;
|
||||||
|
float li = mlp[s * 2 + 1];
|
||||||
|
float nm = fmaxf(m, mi);
|
||||||
|
float corr = __expf(m - nm);
|
||||||
|
float e = __expf(mi - nm);
|
||||||
|
acc = acc * corr + op[s * p.head_dim + d] * e;
|
||||||
|
l = l * corr + li * e;
|
||||||
|
m = nm;
|
||||||
|
}
|
||||||
|
|
||||||
|
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
||||||
|
p.o[(size_t)bh * p.head_dim + d] = __float2bfloat16(acc * inv);
|
||||||
|
}
|
||||||
@@ -0,0 +1,182 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cfloat>
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
#include "attn_mma_utils.cuh"
|
||||||
|
|
||||||
|
using bf16 = __nv_bfloat16;
|
||||||
|
|
||||||
|
// Paged split-KV tensor-core decode via GQA head-packing.
|
||||||
|
// Identical algorithm to attn_decode_split_kv_mma_kernel but reads K/V
|
||||||
|
// directly from the page pool through a page table, eliminating the gather
|
||||||
|
// copy. Each tile (BC=32) fits within a single page (page_size >= 32), so
|
||||||
|
// the page-table lookup happens once per tile for cp.async.
|
||||||
|
//
|
||||||
|
// Optimizations mirror attn_decode_split_kv_mma_kernel:
|
||||||
|
// - Q loaded directly from global into mma A-operand registers (no sQ)
|
||||||
|
// - Double-buffered KV (STAGES=2) for D<=128, single-buffer for D=256
|
||||||
|
// - Predicated cp.async for unified full/partial tile path
|
||||||
|
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
|
||||||
|
__global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16> p) {
|
||||||
|
constexpr int BR = 16;
|
||||||
|
constexpr int KD = HEAD_DIM / 16;
|
||||||
|
constexpr int NC8 = BC / 8;
|
||||||
|
constexpr int KT2 = BC / 16;
|
||||||
|
constexpr int DN8 = HEAD_DIM / 8;
|
||||||
|
constexpr int LD = HEAD_DIM;
|
||||||
|
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
|
||||||
|
constexpr int VEC = 8;
|
||||||
|
constexpr int TOTAL = BC * HEAD_DIM;
|
||||||
|
|
||||||
|
const int lane = threadIdx.x;
|
||||||
|
const int gid = lane >> 2;
|
||||||
|
const int tid4 = lane & 3;
|
||||||
|
|
||||||
|
const int kv_head_idx = blockIdx.x;
|
||||||
|
const int batch = blockIdx.y;
|
||||||
|
const int split = blockIdx.z;
|
||||||
|
const int G = p.q_head / p.kv_head;
|
||||||
|
const int q_head0 = kv_head_idx * G;
|
||||||
|
|
||||||
|
__shared__ __align__(16) bf16 sK[STAGES * BC * LD];
|
||||||
|
__shared__ __align__(16) bf16 sV[STAGES * BC * LD];
|
||||||
|
|
||||||
|
// ---- Load Q directly from global into mma A-operand registers ----
|
||||||
|
const int q_base = (batch * p.q_head + q_head0) * HEAD_DIM;
|
||||||
|
const int qra = gid;
|
||||||
|
const int qrb = gid + 8;
|
||||||
|
const bool va = qra < G, vb = qrb < G;
|
||||||
|
unsigned Qa[KD][4];
|
||||||
|
#pragma unroll
|
||||||
|
for (int kt = 0; kt < KD; kt++) {
|
||||||
|
int c = kt * 16 + tid4 * 2;
|
||||||
|
const unsigned* pau = reinterpret_cast<const unsigned*>(
|
||||||
|
&p.q[q_base + qra * HEAD_DIM + c]);
|
||||||
|
const unsigned* pbu = reinterpret_cast<const unsigned*>(
|
||||||
|
&p.q[q_base + qrb * HEAD_DIM + c]);
|
||||||
|
Qa[kt][0] = va ? pau[0] : 0u;
|
||||||
|
Qa[kt][1] = vb ? pbu[0] : 0u;
|
||||||
|
Qa[kt][2] = va ? pau[4] : 0u;
|
||||||
|
Qa[kt][3] = vb ? pbu[4] : 0u;
|
||||||
|
}
|
||||||
|
|
||||||
|
float Oacc[DN8][4];
|
||||||
|
#pragma unroll
|
||||||
|
for (int j = 0; j < DN8; j++)
|
||||||
|
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
||||||
|
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
||||||
|
|
||||||
|
const int mask_base = batch * p.kv_len;
|
||||||
|
const int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||||
|
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
|
||||||
|
const int ti_begin = split * tiles_per_split;
|
||||||
|
const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
|
||||||
|
const int has_mask = p.use_mask && p.mask;
|
||||||
|
|
||||||
|
// Paged strides (constant for the block)
|
||||||
|
const int64_t page_stride = (int64_t)p.page_size * p.kv_head * HEAD_DIM;
|
||||||
|
const int64_t pos_stride = (int64_t)p.kv_head * HEAD_DIM;
|
||||||
|
const int64_t head_off = (int64_t)kv_head_idx * HEAD_DIM;
|
||||||
|
|
||||||
|
// ---- Load tile lambda: predicated cp.async, paged addressing ----
|
||||||
|
auto load_tile = [&](int ti, int buf) {
|
||||||
|
int kv0 = ti * BC;
|
||||||
|
bf16* dK = sK + buf * BC * LD;
|
||||||
|
bf16* dV = sV + buf * BC * LD;
|
||||||
|
int logical_page = kv0 / p.page_size;
|
||||||
|
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
||||||
|
bool page_valid = (phys_page >= 0);
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = lane * VEC; i < TOTAL; i += 32 * VEC) {
|
||||||
|
int r = i / HEAD_DIM, d = i % HEAD_DIM;
|
||||||
|
int kc = kv0 + r;
|
||||||
|
bool valid = (kc < p.kv_len) && page_valid;
|
||||||
|
int page_off = kc % p.page_size;
|
||||||
|
int64_t gmem_base = (int64_t)phys_page * page_stride
|
||||||
|
+ (int64_t)page_off * pos_stride
|
||||||
|
+ head_off;
|
||||||
|
int off = r * LD + swiz_col(d, r, SWIZ_MASK);
|
||||||
|
cp_async_16_pred(&dK[off], &p.k_cache[gmem_base + d], valid);
|
||||||
|
cp_async_16_pred(&dV[off], &p.v_cache[gmem_base + d], valid);
|
||||||
|
}
|
||||||
|
cp_async_commit();
|
||||||
|
};
|
||||||
|
|
||||||
|
// ---- Prologue: issue first tile load ----
|
||||||
|
if (ti_begin < ti_end) {
|
||||||
|
load_tile(ti_begin, 0);
|
||||||
|
}
|
||||||
|
|
||||||
|
for (int ti = ti_begin; ti < ti_end; ti++) {
|
||||||
|
constexpr int BUF_MASK = (STAGES > 1) ? (STAGES - 1) : 0;
|
||||||
|
int buf = (ti - ti_begin) & BUF_MASK;
|
||||||
|
|
||||||
|
cp_async_wait_group<0>();
|
||||||
|
__syncwarp();
|
||||||
|
if constexpr (STAGES > 1) {
|
||||||
|
if (ti + 1 < ti_end)
|
||||||
|
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
|
||||||
|
}
|
||||||
|
|
||||||
|
const bf16* bK = sK + buf * BC * LD;
|
||||||
|
const bf16* bV = sV + buf * BC * LD;
|
||||||
|
int kv0 = ti * BC;
|
||||||
|
|
||||||
|
float Sacc[NC8][4];
|
||||||
|
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
|
||||||
|
|
||||||
|
#pragma unroll
|
||||||
|
for (int n8 = 0; n8 < NC8; n8++)
|
||||||
|
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||||
|
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||||
|
|
||||||
|
int maxc = p.is_causal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
||||||
|
mma_softmax_tile<NC8, DN8>(kv0, maxc, maxc,
|
||||||
|
mask_base, p.mask, has_mask,
|
||||||
|
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||||
|
|
||||||
|
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc);
|
||||||
|
__syncwarp();
|
||||||
|
|
||||||
|
if constexpr (STAGES == 1) {
|
||||||
|
if (ti + 1 < ti_end)
|
||||||
|
load_tile(ti + 1, 0);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- write UN-normalised partials for this split ----
|
||||||
|
auto split_slot = [&](int h) -> size_t {
|
||||||
|
size_t bh = (size_t)batch * p.q_head + h;
|
||||||
|
return bh * p.num_splits + split;
|
||||||
|
};
|
||||||
|
#pragma unroll
|
||||||
|
for (int dn8 = 0; dn8 < DN8; dn8++) {
|
||||||
|
int d = dn8 * 8 + 2 * tid4;
|
||||||
|
int r0 = gid, r1 = gid + 8;
|
||||||
|
if (r0 < G) {
|
||||||
|
int h = q_head0 + r0;
|
||||||
|
float* op = p.o_part + split_slot(h) * HEAD_DIM;
|
||||||
|
op[d] = Oacc[dn8][0];
|
||||||
|
op[d + 1] = Oacc[dn8][1];
|
||||||
|
}
|
||||||
|
if (r1 < G) {
|
||||||
|
int h = q_head0 + r1;
|
||||||
|
float* op = p.o_part + split_slot(h) * HEAD_DIM;
|
||||||
|
op[d] = Oacc[dn8][2];
|
||||||
|
op[d + 1] = Oacc[dn8][3];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (tid4 == 0) {
|
||||||
|
int r0 = gid, r1 = gid + 8;
|
||||||
|
if (r0 < G) {
|
||||||
|
int h = q_head0 + r0;
|
||||||
|
float* mp = p.ml_part + split_slot(h) * 2;
|
||||||
|
mp[0] = m0; mp[1] = l0;
|
||||||
|
}
|
||||||
|
if (r1 < G) {
|
||||||
|
int h = q_head0 + r1;
|
||||||
|
float* mp = p.ml_part + split_slot(h) * 2;
|
||||||
|
mp[0] = m1; mp[1] = l1;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,79 @@
|
|||||||
|
#include "attn_prefill_split_q.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 q,
|
||||||
|
torch::Tensor k,
|
||||||
|
torch::Tensor v,
|
||||||
|
c10::optional<torch::Tensor> mask,
|
||||||
|
bool is_causal = false,
|
||||||
|
int64_t causal_offset = 0,
|
||||||
|
c10::optional<double> scale = c10::nullopt
|
||||||
|
) {
|
||||||
|
AttentionParams<bf16> p;
|
||||||
|
attn_pack_params(q, k, v, mask, is_causal, causal_offset, scale, p);
|
||||||
|
TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16");
|
||||||
|
|
||||||
|
auto O = torch::empty_like(q);
|
||||||
|
p.o = (bf16*)O.data_ptr();
|
||||||
|
|
||||||
|
switch (p.head_dim) {
|
||||||
|
case 32:
|
||||||
|
dispatch_prefill<32>(p);
|
||||||
|
break;
|
||||||
|
case 64:
|
||||||
|
dispatch_prefill<64>(p);
|
||||||
|
break;
|
||||||
|
case 128:
|
||||||
|
dispatch_prefill<128>(p);
|
||||||
|
break;
|
||||||
|
case 256:
|
||||||
|
dispatch_prefill<256>(p);
|
||||||
|
break;
|
||||||
|
default:
|
||||||
|
TORCH_CHECK(false, "prefill: unsupported head_dim ", p.head_dim,
|
||||||
|
" (supported: 32,64,128,256)");
|
||||||
|
}
|
||||||
|
return O;
|
||||||
|
}
|
||||||
|
|
||||||
|
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||||
|
m.def("attn_prefill", &attn_prefill,
|
||||||
|
py::arg("q"),
|
||||||
|
py::arg("k"),
|
||||||
|
py::arg("v"),
|
||||||
|
py::arg("mask") = py::none(),
|
||||||
|
py::arg("is_causal") = false,
|
||||||
|
py::arg("causal_offset") = 0,
|
||||||
|
py::arg("scale") = py::none(),
|
||||||
|
"GQA prefill (tensor-core mma on sm_80+, scalar fallback)");
|
||||||
|
}
|
||||||
@@ -0,0 +1,140 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cfloat>
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
|
||||||
|
using bf16 = __nv_bfloat16;
|
||||||
|
|
||||||
|
// 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
|
||||||
|
// occupancy high; the S dot product is reduced across the G-lane group with a
|
||||||
|
// short shuffle chain (log2(G) shuffles) instead of a full 32-lane warp reduce.
|
||||||
|
// 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>
|
||||||
|
__device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
|
||||||
|
#pragma unroll
|
||||||
|
for (int o = G / 2; o > 0; o >>= 1)
|
||||||
|
v += __shfl_xor_sync(mask, v, o);
|
||||||
|
return v;
|
||||||
|
}
|
||||||
|
|
||||||
|
// load 8 contiguous bf16 from (16-byte aligned) smem as one float4, unpack to
|
||||||
|
// 8 floats — cuts shared-load instructions 8x vs scalar bf16 loads.
|
||||||
|
__device__ __forceinline__ void ld8(const bf16* p, float* o) {
|
||||||
|
float4 raw = *reinterpret_cast<const float4*>(p);
|
||||||
|
const __nv_bfloat162* h = reinterpret_cast<const __nv_bfloat162*>(&raw);
|
||||||
|
#pragma unroll
|
||||||
|
for (int j = 0; j < 4; j++) {
|
||||||
|
float2 f = __bfloat1622float2(h[j]);
|
||||||
|
o[2 * j] = f.x;
|
||||||
|
o[2 * j + 1] = f.y;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM, int G, int ROWS, int P_BC>
|
||||||
|
__global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
||||||
|
constexpr int DPT = HEAD_DIM / G;
|
||||||
|
|
||||||
|
int q_tile = blockIdx.x;
|
||||||
|
int q_head = blockIdx.y;
|
||||||
|
int batch = blockIdx.z;
|
||||||
|
int gpos = threadIdx.x; // 0..G-1 (which d-chunk)
|
||||||
|
int row = threadIdx.y; // 0..ROWS-1
|
||||||
|
int q_row = q_tile * ROWS + row;
|
||||||
|
|
||||||
|
int kv_head = q_head / (p.q_head / p.kv_head);
|
||||||
|
|
||||||
|
__shared__ __align__(16) bf16 sK[P_BC * HEAD_DIM];
|
||||||
|
__shared__ __align__(16) bf16 sV[P_BC * HEAD_DIM];
|
||||||
|
|
||||||
|
float qreg[DPT];
|
||||||
|
if (q_row < p.q_len) {
|
||||||
|
int q_off = ((batch * p.q_head + q_head) * p.q_len + q_row) * HEAD_DIM + gpos * DPT;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < DPT; i++)
|
||||||
|
qreg[i] = __bfloat162float(p.q[q_off + i]) * p.scale;
|
||||||
|
}
|
||||||
|
|
||||||
|
float m = -FLT_MAX, l = 0.0f;
|
||||||
|
float acc[DPT];
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < DPT; i++)
|
||||||
|
acc[i] = 0.0f;
|
||||||
|
|
||||||
|
int kv_base = ((batch * p.kv_head + kv_head) * p.kv_len) * HEAD_DIM;
|
||||||
|
int tiles = (p.kv_len + P_BC - 1) / P_BC;
|
||||||
|
int tt = G * ROWS;
|
||||||
|
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;
|
||||||
|
unsigned gmask = (G == 32) ? 0xFFFFFFFFu
|
||||||
|
: (((1u << G) - 1u) << (lane_in_warp & ~(G - 1)));
|
||||||
|
|
||||||
|
for (int ti = 0; ti < tiles; ti++) {
|
||||||
|
int kv0 = ti * P_BC;
|
||||||
|
int tlen = min(P_BC, p.kv_len - kv0);
|
||||||
|
|
||||||
|
for (int i = lid; i < tlen * HEAD_DIM; i += tt) {
|
||||||
|
int gidx = kv_base + (kv0 + i / HEAD_DIM) * HEAD_DIM + (i % HEAD_DIM);
|
||||||
|
sK[i] = p.k[gidx];
|
||||||
|
sV[i] = p.v[gidx];
|
||||||
|
}
|
||||||
|
__syncthreads();
|
||||||
|
|
||||||
|
int lim = tlen;
|
||||||
|
if (p.is_causal && q_row < p.q_len) {
|
||||||
|
int ep = q_row + p.causal_offset + 1;
|
||||||
|
if (kv0 >= ep)
|
||||||
|
lim = 0;
|
||||||
|
else if (kv0 + tlen > ep)
|
||||||
|
lim = ep - kv0;
|
||||||
|
}
|
||||||
|
|
||||||
|
for (int s = 0; s < lim; s++) {
|
||||||
|
const bf16* kr = sK + s * HEAD_DIM + gpos * DPT;
|
||||||
|
float part = 0.0f;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < DPT; i += 8) {
|
||||||
|
float k8[8];
|
||||||
|
ld8(kr + i, k8);
|
||||||
|
#pragma unroll
|
||||||
|
for (int j = 0; j < 8; j++)
|
||||||
|
part = fmaf(qreg[i + j], k8[j], part);
|
||||||
|
}
|
||||||
|
float dot = group_reduce_sum<G>(part, gmask);
|
||||||
|
|
||||||
|
if (p.use_mask && p.mask && !p.mask[batch * p.kv_len + kv0 + s])
|
||||||
|
dot = -FLT_MAX;
|
||||||
|
|
||||||
|
float nm = fmaxf(m, dot);
|
||||||
|
float al = __expf(m - nm);
|
||||||
|
float be = __expf(dot - nm);
|
||||||
|
l = l * al + be;
|
||||||
|
|
||||||
|
const bf16* vr = sV + s * HEAD_DIM + gpos * DPT;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < DPT; i += 8) {
|
||||||
|
float v8[8];
|
||||||
|
ld8(vr + i, v8);
|
||||||
|
#pragma unroll
|
||||||
|
for (int j = 0; j < 8; j++)
|
||||||
|
acc[i + j] = fmaf(v8[j], be, acc[i + j] * al);
|
||||||
|
}
|
||||||
|
m = nm;
|
||||||
|
}
|
||||||
|
__syncthreads();
|
||||||
|
}
|
||||||
|
|
||||||
|
if (q_row < p.q_len) {
|
||||||
|
int o_off = ((batch * p.q_head + q_head) * p.q_len + q_row) * HEAD_DIM + gpos * DPT;
|
||||||
|
float rl = (l > 1e-10f) ? (1.0f / l) : 0.0f;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < DPT; i++)
|
||||||
|
p.o[o_off + i] = __float2bfloat16(acc[i] * rl);
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,204 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cfloat>
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
#include "attn_mma_utils.cuh"
|
||||||
|
|
||||||
|
using bf16 = __nv_bfloat16;
|
||||||
|
|
||||||
|
// 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
|
||||||
|
// cores via mma.sync.m16n8k16 (f32 accumulate). Q fragments are loaded once
|
||||||
|
// 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
|
||||||
|
// ahead, so the next tile streams from global memory while the current tile's
|
||||||
|
// 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
|
||||||
|
// (no sQ staging, no prologue barriers); post-multiply scale in float after
|
||||||
|
// 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) {
|
||||||
|
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 lane = threadIdx.x % 32;
|
||||||
|
const int gid = lane >> 2; // 0..7 → rows gid, gid+8
|
||||||
|
const int tid4 = lane & 3; // 0..3
|
||||||
|
const int nthreads = WARPS * 32;
|
||||||
|
|
||||||
|
const int q_head = blockIdx.y;
|
||||||
|
const int batch = blockIdx.z;
|
||||||
|
const int kv_head = q_head / (p.q_head / p.kv_head);
|
||||||
|
const int qrow0 = (blockIdx.x * WARPS + warp) * BR;
|
||||||
|
|
||||||
|
// Static shared memory — sized by template parameters at compile time.
|
||||||
|
// K/V are double-buffered (STAGES=2): the next tile's cp.async load runs
|
||||||
|
// while the current tile's tensor-core math executes, hiding global-load
|
||||||
|
// latency (FA2-style software pipeline). No dynamic smem / carveout opt-in.
|
||||||
|
constexpr int STAGES = 2;
|
||||||
|
__shared__ __align__(16) bf16 sK[STAGES * BC * LD];
|
||||||
|
__shared__ __align__(16) bf16 sV[STAGES * BC * LD];
|
||||||
|
|
||||||
|
// Load the Q fragments straight from global into the mma A-operand layout
|
||||||
|
// (m16n8k16, row-major): no sQ staging area and no serialized per-warp
|
||||||
|
// prologue barriers. Each lane reads exactly the 8 Q elements ldmatrix
|
||||||
|
// would have produced, pre-scaled by the attention scale. Kept resident in
|
||||||
|
// registers across the tile loop.
|
||||||
|
// frag[0]/[2]: row = qrow0 + gid ; frag[1]/[3]: row = qrow0 + gid + 8
|
||||||
|
// frag[0]/[1]: cols kt*16 + tid4*2 + {0,1} ; frag[2]/[3]: + 8
|
||||||
|
const int q_base = ((batch * p.q_head + q_head) * p.q_len) * HEAD_DIM;
|
||||||
|
const int qra = qrow0 + gid;
|
||||||
|
const int qrb = qrow0 + gid + 8;
|
||||||
|
const bool va = qra < p.q_len, vb = qrb < p.q_len;
|
||||||
|
unsigned Qa[KD][4];
|
||||||
|
#pragma unroll
|
||||||
|
for (int kt = 0; kt < KD; kt++) {
|
||||||
|
int c = kt * 16 + tid4 * 2;
|
||||||
|
const unsigned* pau = reinterpret_cast<const unsigned*>(
|
||||||
|
&p.q[q_base + qra * HEAD_DIM + c]);
|
||||||
|
const unsigned* pbu = reinterpret_cast<const unsigned*>(
|
||||||
|
&p.q[q_base + qrb * HEAD_DIM + c]);
|
||||||
|
Qa[kt][0] = va ? pau[0] : 0u;
|
||||||
|
Qa[kt][1] = vb ? pbu[0] : 0u;
|
||||||
|
Qa[kt][2] = va ? pau[4] : 0u;
|
||||||
|
Qa[kt][3] = vb ? pbu[4] : 0u;
|
||||||
|
}
|
||||||
|
|
||||||
|
float Oacc[DN8][4];
|
||||||
|
#pragma unroll
|
||||||
|
for (int j = 0; j < DN8; j++)
|
||||||
|
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
||||||
|
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
||||||
|
|
||||||
|
const int kv_base = ((batch * p.kv_head + kv_head) * p.kv_len) * HEAD_DIM;
|
||||||
|
const int tiles = (p.kv_len + BC - 1) / BC;
|
||||||
|
const int qr0 = qrow0 + gid; // row for c0/c1
|
||||||
|
const int qr1 = qrow0 + gid + 8; // row for c2/c3
|
||||||
|
|
||||||
|
// Causal tile-skip bounds (no-op when is_causal == 0)
|
||||||
|
const int use_skip = p.is_causal;
|
||||||
|
const int max_kv = qrow0 + BR - 1 + p.causal_offset;
|
||||||
|
const int block_max_kv =
|
||||||
|
blockIdx.x * WARPS * BR + WARPS * BR - 1 + p.causal_offset;
|
||||||
|
const int has_mask = p.use_mask && p.mask;
|
||||||
|
const int mb = batch * p.kv_len;
|
||||||
|
|
||||||
|
// 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;
|
||||||
|
if (use_skip) {
|
||||||
|
int bt = block_max_kv / BC;
|
||||||
|
if (bt < t_end) t_end = bt;
|
||||||
|
}
|
||||||
|
|
||||||
|
constexpr int VEC = 8; // bf16 per cp.async unit (16 bytes)
|
||||||
|
constexpr int TOTAL = BC * HEAD_DIM;
|
||||||
|
|
||||||
|
// 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) {
|
||||||
|
int kv0 = ti * BC;
|
||||||
|
bf16* dK = sK + buf * BC * LD;
|
||||||
|
bf16* dV = sV + buf * BC * LD;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = threadIdx.x * VEC; i < TOTAL; i += nthreads * VEC) {
|
||||||
|
int r = i / HEAD_DIM, d = i % HEAD_DIM;
|
||||||
|
int kc = kv0 + r;
|
||||||
|
bool valid = kc < p.kv_len;
|
||||||
|
int off = r * LD + swiz_col(d, r, SWIZ_MASK);
|
||||||
|
cp_async_16_pred(&dK[off], &p.k[kv_base + kc * HEAD_DIM + d], valid);
|
||||||
|
cp_async_16_pred(&dV[off], &p.v[kv_base + kc * HEAD_DIM + d], valid);
|
||||||
|
}
|
||||||
|
cp_async_commit();
|
||||||
|
};
|
||||||
|
|
||||||
|
// Prologue: kick off the first tile's load.
|
||||||
|
load_tile(0, 0);
|
||||||
|
|
||||||
|
for (int ti = 0; ti <= t_end; ti++) {
|
||||||
|
int buf = ti & 1;
|
||||||
|
|
||||||
|
// Wait for the current tile's async copies, then a single barrier: it
|
||||||
|
// 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>();
|
||||||
|
__syncthreads();
|
||||||
|
if (ti < t_end) load_tile(ti + 1, (ti + 1) & 1);
|
||||||
|
|
||||||
|
const bf16* bK = sK + buf * BC * LD;
|
||||||
|
const bf16* bV = sV + buf * BC * LD;
|
||||||
|
int kv0 = ti * BC;
|
||||||
|
|
||||||
|
// Warp-level causal skip
|
||||||
|
if (!use_skip || kv0 <= max_kv) {
|
||||||
|
|
||||||
|
// S = Q @ K^T + scale + online softmax + O += P @ V
|
||||||
|
float Sacc[NC8][4];
|
||||||
|
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)
|
||||||
|
#pragma unroll
|
||||||
|
for (int n8 = 0; n8 < NC8; n8++)
|
||||||
|
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||||
|
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||||
|
|
||||||
|
int maxc0 = p.is_causal ? min(p.kv_len, qr0 + p.causal_offset + 1)
|
||||||
|
: p.kv_len;
|
||||||
|
int maxc1 = p.is_causal ? min(p.kv_len, qr1 + p.causal_offset + 1)
|
||||||
|
: p.kv_len;
|
||||||
|
mma_softmax_tile<NC8, DN8>(kv0, maxc0, maxc1,
|
||||||
|
mb, p.mask, has_mask,
|
||||||
|
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||||
|
|
||||||
|
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc);
|
||||||
|
} // if active (warp-level causal skip)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- write output ---- (packed bf16x2 stores: one 32-bit STG per pair,
|
||||||
|
// halves store count and removes the uncoalesced scalar-store penalty)
|
||||||
|
float rl0 = (l0 > 1e-20f) ? (1.0f / l0) : 0.0f;
|
||||||
|
float rl1 = (l1 > 1e-20f) ? (1.0f / l1) : 0.0f;
|
||||||
|
const int o_base = ((batch * p.q_head + q_head) * p.q_len) * HEAD_DIM;
|
||||||
|
#pragma unroll
|
||||||
|
for (int dn8 = 0; dn8 < DN8; dn8++) {
|
||||||
|
int d = dn8 * 8 + 2 * tid4;
|
||||||
|
if (qr0 < p.q_len) {
|
||||||
|
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][0] * rl0,
|
||||||
|
Oacc[dn8][1] * rl0);
|
||||||
|
*reinterpret_cast<__nv_bfloat162*>(&p.o[o_base + qr0 * HEAD_DIM + d]) = v;
|
||||||
|
}
|
||||||
|
if (qr1 < p.q_len) {
|
||||||
|
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1,
|
||||||
|
Oacc[dn8][3] * rl1);
|
||||||
|
*reinterpret_cast<__nv_bfloat162*>(&p.o[o_base + qr1 * HEAD_DIM + d]) = v;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,199 @@
|
|||||||
|
/*
|
||||||
|
Pure-C test:
|
||||||
|
nvcc -I csrc -arch=sm_89 -O3 \
|
||||||
|
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
|
||||||
|
csrc/tests/attn_decode_test.cu -o test && ./test
|
||||||
|
*/
|
||||||
|
|
||||||
|
#include "test_utils.cuh"
|
||||||
|
#include "../kernels/attn_decode_split_kv.cuh"
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
#include "../kernels/attn_decode_split_kv_mma.cuh"
|
||||||
|
#endif
|
||||||
|
|
||||||
|
// Split-K scratch (torch-free): the production launcher allocates these from
|
||||||
|
// torch; here we pass pre-allocated device buffers so the bench loop doesn't
|
||||||
|
// pay a cudaMalloc per iteration. Size for the maximum split count (32).
|
||||||
|
struct DecodeScratch {
|
||||||
|
float* o_part = nullptr;
|
||||||
|
float* ml_part = nullptr;
|
||||||
|
};
|
||||||
|
|
||||||
|
// Launch the production decode path (tensor-core head-packing MMA on sm_80+,
|
||||||
|
// scalar fallback otherwise), mirroring dispatch_decode() in attn_decode.cu.
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
static bool decode_use_mma(const AttentionParams<bf16>& p) {
|
||||||
|
int G = p.q_head / p.kv_head;
|
||||||
|
return !p.use_mask && G > 1 && G <= 16;
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
|
||||||
|
static void launch_mma_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
||||||
|
int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||||
|
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
||||||
|
p.o_part = sc.o_part;
|
||||||
|
p.ml_part = sc.ml_part;
|
||||||
|
|
||||||
|
attn_decode_split_kv_mma_kernel<HEAD_DIM, BC, STAGES>
|
||||||
|
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
|
||||||
|
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||||
|
}
|
||||||
|
#endif
|
||||||
|
|
||||||
|
static void launch_scalar_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
||||||
|
int gs = p.q_head / p.kv_head;
|
||||||
|
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
||||||
|
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||||
|
p.o_part = sc.o_part;
|
||||||
|
p.ml_part = sc.ml_part;
|
||||||
|
|
||||||
|
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
||||||
|
attn_decode_split_kv_kernel<<<dim3(p.batch * p.kv_head, 1, p.num_splits), dim3(32, gs), smem>>>(p);
|
||||||
|
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static void dispatch_decode_t(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
if (decode_use_mma(p)) { launch_mma_decode<HEAD_DIM, 32>(p, sc); return; }
|
||||||
|
#endif
|
||||||
|
launch_scalar_decode(p, sc);
|
||||||
|
}
|
||||||
|
|
||||||
|
static void dispatch_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
||||||
|
dispatch_by_head_dim(p.head_dim, [&]<int D>() { dispatch_decode_t<D>(p, sc); });
|
||||||
|
}
|
||||||
|
|
||||||
|
// Warmed-up, CUDA-event timed sweep over the production decode MMA path.
|
||||||
|
static void bench() {
|
||||||
|
const int cfgs[][5] = {
|
||||||
|
{1, 32, 4, 512, 128}, // B, Hq, Hk, kv_len, D
|
||||||
|
{1, 32, 4, 1024, 128},
|
||||||
|
{1, 32, 4, 2048, 128},
|
||||||
|
{1, 32, 4, 4096, 128},
|
||||||
|
{16, 32, 4, 2048, 128},
|
||||||
|
{32, 32, 4, 1024, 128},
|
||||||
|
};
|
||||||
|
const int WARMUP = 10, ITERS = 100;
|
||||||
|
printf("\n===== DECODE BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
|
||||||
|
print_bench_header();
|
||||||
|
|
||||||
|
for (int ci = 0; ci < 6; ci++) {
|
||||||
|
int B = cfgs[ci][0], Hq = cfgs[ci][1], Hk = cfgs[ci][2];
|
||||||
|
int sl = cfgs[ci][3], D = cfgs[ci][4];
|
||||||
|
size_t nQ = (size_t)B * Hq * D;
|
||||||
|
size_t nKV = (size_t)B * Hk * sl * D;
|
||||||
|
|
||||||
|
bf16 *dQ, *dK, *dV, *dO;
|
||||||
|
cudaMalloc(&dQ, nQ*2); cudaMalloc(&dK, nKV*2);
|
||||||
|
cudaMalloc(&dV, nKV*2); cudaMalloc(&dO, nQ*2);
|
||||||
|
size_t big = nQ > nKV ? nQ : nKV; bf16* tmp = new bf16[big];
|
||||||
|
for (size_t i = 0; i < nQ; i++) tmp[i] = f2bf(randf());
|
||||||
|
cudaMemcpy(dQ, tmp, nQ*2, cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
|
||||||
|
cudaMemcpy(dK, tmp, nKV*2, cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
|
||||||
|
cudaMemcpy(dV, tmp, nKV*2, cudaMemcpyHostToDevice);
|
||||||
|
delete[] tmp;
|
||||||
|
|
||||||
|
AttentionParams<bf16> p;
|
||||||
|
p.batch = B; p.q_head = Hq; p.kv_head = Hk; p.q_len = 1; p.kv_len = sl;
|
||||||
|
p.head_dim = D; p.use_mask = 0; p.is_causal = 0; p.causal_offset = 0;
|
||||||
|
p.scale = 1.0f / sqrtf((float)D);
|
||||||
|
p.q = dQ; p.k = dK; p.v = dV; p.mask = nullptr; p.o = dO;
|
||||||
|
|
||||||
|
DecodeScratch sc;
|
||||||
|
cudaMalloc(&sc.o_part, (size_t)B*Hq*32*D*sizeof(float));
|
||||||
|
cudaMalloc(&sc.ml_part, (size_t)B*Hq*32*2*sizeof(float));
|
||||||
|
|
||||||
|
auto launch = [&]() { dispatch_decode(p, sc); };
|
||||||
|
double flops = 4.0 * B * Hq * (double)sl * D;
|
||||||
|
double bytes = 2.0 * (2.0 * nKV * sizeof(bf16));
|
||||||
|
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops, bytes);
|
||||||
|
|
||||||
|
char cfg[64];
|
||||||
|
snprintf(cfg, sizeof(cfg),
|
||||||
|
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
|
||||||
|
B, Hq, Hk, 1, sl, D, 0);
|
||||||
|
print_bench_row(cfg, r);
|
||||||
|
|
||||||
|
cudaFree(dQ); cudaFree(dK); cudaFree(dV); cudaFree(dO);
|
||||||
|
cudaFree(sc.o_part); cudaFree(sc.ml_part);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
int main() {
|
||||||
|
const int configs[][5] = {
|
||||||
|
{1, 2, 1, 64, 32}, // B,Hq,Hk,seq_len,D
|
||||||
|
{1, 32, 4, 512, 128},
|
||||||
|
{1, 32, 4, 1024, 128},
|
||||||
|
};
|
||||||
|
int n_cfgs = sizeof(configs) / sizeof(configs[0]);
|
||||||
|
|
||||||
|
for (int ci = 0; ci < n_cfgs; ci++) {
|
||||||
|
int B = configs[ci][0], Hq = configs[ci][1], Hk = configs[ci][2];
|
||||||
|
int sl = configs[ci][3], D = configs[ci][4], gs = Hq / Hk;
|
||||||
|
printf("=== B=%d Hq=%d Hk=%d seq=%d D=%d gs=%d ===\n", B,Hq,Hk,sl,D,gs);
|
||||||
|
|
||||||
|
size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D;
|
||||||
|
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
||||||
|
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
||||||
|
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
||||||
|
|
||||||
|
bool* hMask=new bool[B*sl];
|
||||||
|
for (int i=0;i<B*sl;i++) hMask[i]=true;
|
||||||
|
|
||||||
|
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||||
|
bool* dMask;
|
||||||
|
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||||
|
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||||
|
cudaMalloc(&dMask,B*sl);
|
||||||
|
|
||||||
|
tmp=new bf16[max(nQ,nKV)];
|
||||||
|
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
||||||
|
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
||||||
|
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
||||||
|
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||||
|
cudaMemcpy(dMask,hMask,B*sl,cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
AttentionParams<bf16> p;
|
||||||
|
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D;
|
||||||
|
p.use_mask=0; p.is_causal=0; p.causal_offset=0;
|
||||||
|
p.scale=1.0f/sqrtf((float)D);
|
||||||
|
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||||
|
|
||||||
|
// Split-K scratch (max 32 splits), sized for the production MMA path.
|
||||||
|
DecodeScratch sc;
|
||||||
|
cudaMalloc(&sc.o_part, (size_t)B*Hq*32*D*sizeof(float));
|
||||||
|
cudaMalloc(&sc.ml_part, (size_t)B*Hq*32*2*sizeof(float));
|
||||||
|
|
||||||
|
double t0=now_ms();
|
||||||
|
dispatch_decode(p, sc);
|
||||||
|
cudaDeviceSynchronize();
|
||||||
|
double kms=now_ms()-t0;
|
||||||
|
cudaError_t err=cudaGetLastError();
|
||||||
|
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
||||||
|
|
||||||
|
bf16* hOut=new bf16[nQ];
|
||||||
|
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
||||||
|
|
||||||
|
float* ref=new float[nQ];
|
||||||
|
cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, 0, 0);
|
||||||
|
|
||||||
|
float max_err=0;
|
||||||
|
for (size_t i=0;i<nQ;i++){
|
||||||
|
float d=fabsf(bf2f(hOut[i])-ref[i]);
|
||||||
|
if(d>max_err) max_err=d;
|
||||||
|
}
|
||||||
|
printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err);
|
||||||
|
|
||||||
|
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask);
|
||||||
|
cudaFree(sc.o_part);cudaFree(sc.ml_part);
|
||||||
|
delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp;
|
||||||
|
}
|
||||||
|
printf("All tests passed!\n");
|
||||||
|
bench();
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
@@ -0,0 +1,330 @@
|
|||||||
|
// Compile:
|
||||||
|
// nvcc -I csrc -arch=sm_89 -O3 --use_fast_math --ptxas-options=-O3 \
|
||||||
|
// --extra-device-vectorization csrc/tests/attn_paged_decode_test.cu \
|
||||||
|
// -o /tmp/test_paged && /tmp/test_paged
|
||||||
|
|
||||||
|
#include <cstring>
|
||||||
|
#include "test_utils.cuh"
|
||||||
|
#include "../kernels/attn_paged_decode_split_kv.cuh"
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
#include "../kernels/attn_paged_decode_split_kv_mma.cuh"
|
||||||
|
#endif
|
||||||
|
|
||||||
|
// Copy contiguous K/V from page pool (reference gather)
|
||||||
|
static void gather_kv_cpu(
|
||||||
|
const bf16* h_k_pool, const bf16* h_v_pool,
|
||||||
|
const int64_t* h_pt, int B, int Hkv, int kv_len,
|
||||||
|
int page_size, int head_dim,
|
||||||
|
bf16* h_k, bf16* h_v)
|
||||||
|
{
|
||||||
|
int max_pages = (kv_len + page_size - 1) / page_size;
|
||||||
|
size_t page_stride = (size_t)page_size * Hkv * head_dim;
|
||||||
|
for (int b = 0; b < B; b++) {
|
||||||
|
for (int pos = 0; pos < kv_len; pos++) {
|
||||||
|
int log_pg = pos / page_size;
|
||||||
|
int pg_off = pos % page_size;
|
||||||
|
int phys = (int)h_pt[b * max_pages + log_pg];
|
||||||
|
for (int h = 0; h < Hkv; h++) {
|
||||||
|
size_t src_base = (size_t)phys * page_stride
|
||||||
|
+ (size_t)pg_off * Hkv * head_dim
|
||||||
|
+ h * head_dim;
|
||||||
|
size_t dst_base = ((size_t)b * Hkv + h) * kv_len * head_dim + (size_t)pos * head_dim;
|
||||||
|
memcpy(h_k + dst_base, h_k_pool + src_base, head_dim * sizeof(bf16));
|
||||||
|
memcpy(h_v + dst_base, h_v_pool + src_base, head_dim * sizeof(bf16));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static void launch_paged_decode(PagedAttentionParams<bf16, float>& p) {
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
int G_check = p.q_head / p.kv_head;
|
||||||
|
bool use_mma = !p.use_mask && G_check >= 1 && G_check <= 16 && p.page_size >= 32;
|
||||||
|
if (use_mma) {
|
||||||
|
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
|
||||||
|
int tiles_total = (p.kv_len + 32 - 1) / 32;
|
||||||
|
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
||||||
|
paged_attn_decode_split_kv_mma_kernel<HEAD_DIM, 32, STAGES>
|
||||||
|
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
|
||||||
|
} else
|
||||||
|
#endif
|
||||||
|
{
|
||||||
|
int group_sz = p.q_head / p.kv_head;
|
||||||
|
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
||||||
|
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||||
|
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
|
||||||
|
paged_attn_decode_split_kv_kernel<<<
|
||||||
|
dim3(p.batch * p.kv_head, 1, p.num_splits),
|
||||||
|
dim3(32, group_sz), smem>>>(p);
|
||||||
|
}
|
||||||
|
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed) {
|
||||||
|
printf("B=%d Hq=%d Hkv=%d kv_len=%d page_sz=%d head_dim=%d ... ", B, Hq, Hkv, kv_len, page_size, HEAD_DIM);
|
||||||
|
fflush(stdout);
|
||||||
|
|
||||||
|
int max_pages = (kv_len + page_size - 1) / page_size;
|
||||||
|
int n_phys_pages = B * max_pages;
|
||||||
|
|
||||||
|
size_t sz_q = (size_t)B * Hq * 1 * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_o = sz_q;
|
||||||
|
size_t sz_kv = (size_t)n_phys_pages * page_size * Hkv * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_pt = (size_t)B * max_pages * sizeof(int64_t);
|
||||||
|
int max_splits = 32;
|
||||||
|
size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float);
|
||||||
|
size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float);
|
||||||
|
|
||||||
|
bf16 *d_q, *d_o_paged, *d_o_ref;
|
||||||
|
bf16 *d_k_pool, *d_v_pool;
|
||||||
|
int64_t* d_pt;
|
||||||
|
float *d_op, *d_ml;
|
||||||
|
|
||||||
|
cudaMalloc(&d_q, sz_q);
|
||||||
|
cudaMalloc(&d_o_paged, sz_o);
|
||||||
|
cudaMalloc(&d_o_ref, sz_o);
|
||||||
|
cudaMalloc(&d_k_pool, sz_kv);
|
||||||
|
cudaMalloc(&d_v_pool, sz_kv);
|
||||||
|
cudaMalloc(&d_pt, sz_pt);
|
||||||
|
cudaMalloc(&d_op, sz_op);
|
||||||
|
cudaMalloc(&d_ml, sz_ml);
|
||||||
|
|
||||||
|
srand(seed);
|
||||||
|
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
|
||||||
|
|
||||||
|
bf16* h_q = (bf16*)malloc(sz_q);
|
||||||
|
for (int i = 0; i < B * Hq * HEAD_DIM; i++)
|
||||||
|
h_q[i] = __float2bfloat16(rnd());
|
||||||
|
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
bf16* h_k_pool = (bf16*)malloc(sz_kv);
|
||||||
|
bf16* h_v_pool = (bf16*)malloc(sz_kv);
|
||||||
|
size_t ps = (size_t)page_size * Hkv * HEAD_DIM;
|
||||||
|
for (int pg = 0; pg < n_phys_pages; pg++) {
|
||||||
|
for (int off = 0; off < page_size; off++) {
|
||||||
|
for (int h = 0; h < Hkv; h++) {
|
||||||
|
for (int d = 0; d < HEAD_DIM; d++) {
|
||||||
|
float v = sinf((float)(pg * 7919 + off * 1049 + h * 331 + d));
|
||||||
|
size_t idx = (size_t)pg * ps + (size_t)off * Hkv * HEAD_DIM + h * HEAD_DIM + d;
|
||||||
|
h_k_pool[idx] = __float2bfloat16(v);
|
||||||
|
h_v_pool[idx] = __float2bfloat16(v * 0.3f);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
int64_t* h_pt = (int64_t*)malloc(sz_pt);
|
||||||
|
int next_pg = 0;
|
||||||
|
for (int b = 0; b < B; b++)
|
||||||
|
for (int p = 0; p < max_pages; p++)
|
||||||
|
h_pt[b * max_pages + p] = next_pg++;
|
||||||
|
cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
bf16* h_k_cont = (bf16*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(bf16));
|
||||||
|
bf16* h_v_cont = (bf16*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(bf16));
|
||||||
|
gather_kv_cpu(h_k_pool, h_v_pool, h_pt, B, Hkv, kv_len, page_size, HEAD_DIM, h_k_cont, h_v_cont);
|
||||||
|
|
||||||
|
float* h_q_f = (float*)malloc((size_t)B * Hq * HEAD_DIM * sizeof(float));
|
||||||
|
float* h_k_f = (float*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(float));
|
||||||
|
float* h_v_f = (float*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(float));
|
||||||
|
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
|
||||||
|
for (int i = 0; i < B * kv_len * Hkv * HEAD_DIM; i++) {
|
||||||
|
h_k_f[i] = bf2f(h_k_cont[i]);
|
||||||
|
h_v_f[i] = bf2f(h_v_cont[i]);
|
||||||
|
}
|
||||||
|
|
||||||
|
float* h_o_ref = (float*)calloc(B * Hq * HEAD_DIM, sizeof(float));
|
||||||
|
cpu_attention_ref(h_q_f, h_k_f, h_v_f, nullptr, h_o_ref, B, Hq, Hkv, 1, kv_len, HEAD_DIM, 0, 0);
|
||||||
|
|
||||||
|
float scale_val = 1.0f / sqrtf((float)HEAD_DIM);
|
||||||
|
PagedAttentionParams<bf16, float> p;
|
||||||
|
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.q_len = 1;
|
||||||
|
p.kv_len = kv_len; p.head_dim = HEAD_DIM;
|
||||||
|
p.use_mask = 0; p.is_causal = 0; p.causal_offset = 0;
|
||||||
|
p.num_splits = 1; p.scale = scale_val;
|
||||||
|
p.page_size = page_size; p.max_pages = max_pages;
|
||||||
|
p.page_table = d_pt;
|
||||||
|
p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||||
|
p.q = d_q; p.mask = nullptr; p.o = d_o_paged;
|
||||||
|
p.o_part = d_op; p.ml_part = d_ml;
|
||||||
|
|
||||||
|
launch_paged_decode<HEAD_DIM>(p);
|
||||||
|
cudaDeviceSynchronize();
|
||||||
|
|
||||||
|
bf16* h_o_bf16 = (bf16*)malloc(sz_o);
|
||||||
|
cudaMemcpy(h_o_bf16, d_o_paged, sz_o, cudaMemcpyDeviceToHost);
|
||||||
|
float* h_o_paged = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
|
||||||
|
for (int i = 0; i < B * Hq * HEAD_DIM; i++)
|
||||||
|
h_o_paged[i] = __bfloat162float(h_o_bf16[i]);
|
||||||
|
|
||||||
|
float max_err = 0.0f;
|
||||||
|
int bad_idx = -1;
|
||||||
|
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
|
||||||
|
float e = fabsf(h_o_paged[i] - h_o_ref[i]);
|
||||||
|
if (e > max_err) { max_err = e; bad_idx = i; }
|
||||||
|
}
|
||||||
|
|
||||||
|
bool pass = max_err < 0.02f;
|
||||||
|
|
||||||
|
if (pass) {
|
||||||
|
printf("PASS (max_abs_err=%.4e)\n", max_err);
|
||||||
|
} else {
|
||||||
|
int b = bad_idx / (Hq * HEAD_DIM);
|
||||||
|
int h = (bad_idx / HEAD_DIM) % Hq;
|
||||||
|
int d = bad_idx % HEAD_DIM;
|
||||||
|
printf("FAIL (max_abs_err=%.4e at [%d,%d,%d]: ref=%.4f got=%.4f)\n",
|
||||||
|
max_err, b, h, d, h_o_ref[bad_idx], h_o_paged[bad_idx]);
|
||||||
|
printf(" ref[0..7]:");
|
||||||
|
for (int i = 0; i < 8 && i < HEAD_DIM; i++)
|
||||||
|
printf(" %.4f", h_o_ref[i]);
|
||||||
|
printf("\n got[0..7]:");
|
||||||
|
for (int i = 0; i < 8 && i < HEAD_DIM; i++)
|
||||||
|
printf(" %.4f", h_o_paged[i]);
|
||||||
|
printf("\n");
|
||||||
|
}
|
||||||
|
|
||||||
|
free(h_q); free(h_k_pool); free(h_v_pool); free(h_pt);
|
||||||
|
free(h_k_cont); free(h_v_cont);
|
||||||
|
free(h_q_f); free(h_k_f); free(h_v_f);
|
||||||
|
free(h_o_ref); free(h_o_bf16); free(h_o_paged);
|
||||||
|
cudaFree(d_q); cudaFree(d_o_paged); cudaFree(d_o_ref);
|
||||||
|
cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_pt);
|
||||||
|
cudaFree(d_op); cudaFree(d_ml);
|
||||||
|
|
||||||
|
return pass ? 0 : 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
struct TestCase {
|
||||||
|
int head_dim;
|
||||||
|
int B, Hq, Hkv, kv_len, page_size, seed;
|
||||||
|
};
|
||||||
|
|
||||||
|
static const TestCase TESTS[] = {
|
||||||
|
{128, 1, 1, 1, 8, 128, 1},
|
||||||
|
{128, 1, 4, 4, 128, 128, 2},
|
||||||
|
{128, 2, 4, 4, 256, 128, 3},
|
||||||
|
{128, 1, 4, 1, 64, 64, 4},
|
||||||
|
{128, 1, 8, 2, 64, 128, 5},
|
||||||
|
{128, 2, 16, 4, 128, 128, 6},
|
||||||
|
{64, 1, 4, 2, 32, 128, 7},
|
||||||
|
{256, 1, 2, 1, 16, 128, 8},
|
||||||
|
{32, 1, 4, 2, 32, 64, 9},
|
||||||
|
{128, 3, 8, 2, 256, 128, 10},
|
||||||
|
{128, 2, 32, 8, 512, 128, 11},
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
{128, 1, 16, 2, 256, 128, 12},
|
||||||
|
{128, 2, 32, 4, 512, 128, 13},
|
||||||
|
#endif
|
||||||
|
};
|
||||||
|
|
||||||
|
static int dispatch_test(const TestCase& tc) {
|
||||||
|
bool matched = false;
|
||||||
|
int r = 0;
|
||||||
|
dispatch_by_head_dim(tc.head_dim, [&]<int D>() {
|
||||||
|
matched = true;
|
||||||
|
r = run_test<D>(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.seed);
|
||||||
|
});
|
||||||
|
return matched ? r : 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Warmed-up, CUDA-event timed sweep over paged decode configs.
|
||||||
|
// Bytes = K + V read through page table (B*Hk*kv*D each), bf16.
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) {
|
||||||
|
int max_pages = (kv_len + page_size - 1) / page_size;
|
||||||
|
int n_phys_pages = B * max_pages;
|
||||||
|
|
||||||
|
size_t sz_q = (size_t)B * Hq * 1 * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_kv = (size_t)n_phys_pages * page_size * Hkv * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_pt = (size_t)B * max_pages * sizeof(int64_t);
|
||||||
|
int max_splits = 32;
|
||||||
|
size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float);
|
||||||
|
size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float);
|
||||||
|
|
||||||
|
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
|
||||||
|
int64_t* d_pt;
|
||||||
|
float *d_op, *d_ml;
|
||||||
|
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
|
||||||
|
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
|
||||||
|
cudaMalloc(&d_pt, sz_pt);
|
||||||
|
cudaMalloc(&d_op, sz_op); cudaMalloc(&d_ml, sz_ml);
|
||||||
|
|
||||||
|
bf16* tmp = (bf16*)malloc(sz_kv > sz_q ? sz_kv : sz_q);
|
||||||
|
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) tmp[i] = f2bf(randf());
|
||||||
|
cudaMemcpy(d_q, tmp, sz_q, cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) tmp[i] = f2bf(randf());
|
||||||
|
cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
int64_t* h_pt = (int64_t*)malloc(sz_pt);
|
||||||
|
int next_pg = 0;
|
||||||
|
for (int b = 0; b < B; b++)
|
||||||
|
for (int p = 0; p < max_pages; p++)
|
||||||
|
h_pt[b * max_pages + p] = next_pg++;
|
||||||
|
cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice);
|
||||||
|
free(h_pt);
|
||||||
|
|
||||||
|
float scale_val = 1.0f / sqrtf((float)HEAD_DIM);
|
||||||
|
PagedAttentionParams<bf16, float> pa;
|
||||||
|
pa.batch = B; pa.q_head = Hq; pa.kv_head = Hkv; pa.q_len = 1;
|
||||||
|
pa.kv_len = kv_len; pa.head_dim = HEAD_DIM;
|
||||||
|
pa.use_mask = 0; pa.is_causal = 0; pa.causal_offset = 0;
|
||||||
|
pa.num_splits = 1; pa.scale = scale_val;
|
||||||
|
pa.page_size = page_size; pa.max_pages = max_pages;
|
||||||
|
pa.page_table = d_pt;
|
||||||
|
pa.k_cache = d_k_pool; pa.v_cache = d_v_pool;
|
||||||
|
pa.q = d_q; pa.mask = nullptr; pa.o = d_o;
|
||||||
|
pa.o_part = d_op; pa.ml_part = d_ml;
|
||||||
|
|
||||||
|
const int WARMUP = 10, ITERS = 100;
|
||||||
|
auto launch = [&]() { launch_paged_decode<HEAD_DIM>(pa); };
|
||||||
|
double flops = 4.0 * B * Hq * (double)kv_len * HEAD_DIM;
|
||||||
|
size_t nKV = (size_t)B * Hkv * kv_len * HEAD_DIM;
|
||||||
|
double bytes = 2.0 * (2.0 * nKV * sizeof(bf16));
|
||||||
|
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops, bytes);
|
||||||
|
|
||||||
|
char cfg[64];
|
||||||
|
snprintf(cfg, sizeof(cfg),
|
||||||
|
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d page=%3d",
|
||||||
|
B, Hq, Hkv, 1, kv_len, HEAD_DIM, page_size);
|
||||||
|
print_bench_row(cfg, r);
|
||||||
|
|
||||||
|
free(tmp);
|
||||||
|
cudaFree(d_q); cudaFree(d_o);
|
||||||
|
cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_pt);
|
||||||
|
cudaFree(d_op); cudaFree(d_ml);
|
||||||
|
}
|
||||||
|
|
||||||
|
static void bench() {
|
||||||
|
printf("\n===== PAGED DECODE BENCH =====\n");
|
||||||
|
print_bench_header();
|
||||||
|
bench_config<128>(1, 32, 4, 512, 128);
|
||||||
|
bench_config<128>(1, 32, 4, 1024, 128);
|
||||||
|
bench_config<128>(1, 32, 4, 2048, 128);
|
||||||
|
bench_config<128>(1, 32, 4, 4096, 128);
|
||||||
|
bench_config<128>(16, 32, 4, 2048, 128);
|
||||||
|
bench_config<128>(32, 32, 4, 1024, 128);
|
||||||
|
}
|
||||||
|
|
||||||
|
int main() {
|
||||||
|
int n = sizeof(TESTS) / sizeof(TESTS[0]);
|
||||||
|
int fail = 0;
|
||||||
|
printf("=== Paged Decode vs CPU reference (%d cases) ===\n\n", n);
|
||||||
|
|
||||||
|
for (int i = 0; i < n; i++) {
|
||||||
|
fail += dispatch_test(TESTS[i]);
|
||||||
|
if (fail) break;
|
||||||
|
}
|
||||||
|
|
||||||
|
if (fail) {
|
||||||
|
printf("\nFAILED (%d/%d tests failed)\n", fail, n);
|
||||||
|
return fail;
|
||||||
|
}
|
||||||
|
printf("\nAll %d tests passed!\n", n);
|
||||||
|
bench();
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
@@ -0,0 +1,176 @@
|
|||||||
|
/*
|
||||||
|
Pure-C test:
|
||||||
|
nvcc -I csrc -arch=sm_89 -O3 \
|
||||||
|
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
|
||||||
|
csrc/tests/attn_prefill_test.cu -o test && ./test
|
||||||
|
*/
|
||||||
|
|
||||||
|
#include "test_utils.cuh"
|
||||||
|
#include "../kernels/attn_prefill_split_q.cuh"
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
#include "../kernels/attn_prefill_split_q_mma.cuh"
|
||||||
|
#endif
|
||||||
|
|
||||||
|
// Launch the production prefill path (tensor-core MMA on sm_80+, else the
|
||||||
|
// scalar fallback), mirroring dispatch_prefill() in attn_prefill.cu.
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static void launch_prefill(AttentionParams<bf16>& p) {
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
constexpr int WARPS = 4, BR = 16;
|
||||||
|
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
|
||||||
|
dim3 grid((p.q_len + BR * WARPS - 1) / (BR * WARPS), p.q_head, p.batch);
|
||||||
|
dim3 block(WARPS * 32, 1, 1);
|
||||||
|
attn_prefill_split_q_mma_kernel<HEAD_DIM, WARPS, BC><<<grid, block>>>(p);
|
||||||
|
#else
|
||||||
|
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
||||||
|
dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
|
||||||
|
dim3 block(G, ROWS, 1);
|
||||||
|
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC><<<grid, block>>>(p);
|
||||||
|
#endif
|
||||||
|
}
|
||||||
|
|
||||||
|
static void dispatch_prefill(AttentionParams<bf16>& p) {
|
||||||
|
switch (p.head_dim) {
|
||||||
|
case 64: launch_prefill<64>(p); break;
|
||||||
|
case 128: launch_prefill<128>(p); break;
|
||||||
|
default: printf("bench: unsupported D=%d\n", p.head_dim);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Warmed-up, CUDA-event timed throughput sweep over the production MMA path.
|
||||||
|
// Reports per-call latency and effective tensor-core TFLOP/s (2 matmuls:
|
||||||
|
// QK^T and P@V, each 2*B*Hq*ql*kl*D flops; halved for causal).
|
||||||
|
static void bench() {
|
||||||
|
const int cfgs[][7] = {
|
||||||
|
{1,32,4,512,512,128,0},
|
||||||
|
{1,32,4,1024,1024,128,0},
|
||||||
|
{1,32,4,2048,2048,128,0},
|
||||||
|
{1,32,4,2048,2048,128,1},
|
||||||
|
{4,32,4,2048,2048,128,1},
|
||||||
|
{1,32,4,4096,4096,128,1},
|
||||||
|
};
|
||||||
|
int n = sizeof(cfgs)/sizeof(cfgs[0]);
|
||||||
|
const int WARMUP = 10, ITERS = 50;
|
||||||
|
printf("\n===== PREFILL BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
|
||||||
|
printf("%-46s | %10s | %10s | %10s\n",
|
||||||
|
"config", "latency", "bandwidth", "throughput");
|
||||||
|
printf("---------------------------------------------------------------"
|
||||||
|
"----------------------------\n");
|
||||||
|
|
||||||
|
for (int ci = 0; ci < n; ci++) {
|
||||||
|
int B=cfgs[ci][0], Hq=cfgs[ci][1], Hk=cfgs[ci][2];
|
||||||
|
int ql=cfgs[ci][3], kl=cfgs[ci][4], D=cfgs[ci][5], causal=cfgs[ci][6];
|
||||||
|
size_t nQ=(size_t)B*Hq*ql*D, nKV=(size_t)B*Hk*kl*D;
|
||||||
|
|
||||||
|
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||||
|
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||||
|
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||||
|
size_t big = nQ>nKV?nQ:nKV; tmp=new bf16[big];
|
||||||
|
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(randf());
|
||||||
|
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
|
||||||
|
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
|
||||||
|
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
AttentionParams<bf16> p;
|
||||||
|
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
|
||||||
|
p.use_mask=0; p.is_causal=causal; p.causal_offset=0;
|
||||||
|
p.scale=1.0f/sqrtf((float)D);
|
||||||
|
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||||
|
|
||||||
|
for (int i=0;i<WARMUP;i++) dispatch_prefill(p);
|
||||||
|
cudaDeviceSynchronize();
|
||||||
|
cudaError_t err=cudaGetLastError();
|
||||||
|
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return;}
|
||||||
|
|
||||||
|
cudaEvent_t s,e; cudaEventCreate(&s); cudaEventCreate(&e);
|
||||||
|
cudaEventRecord(s);
|
||||||
|
for (int i=0;i<ITERS;i++) dispatch_prefill(p);
|
||||||
|
cudaEventRecord(e); cudaEventSynchronize(e);
|
||||||
|
float ms=0; cudaEventElapsedTime(&ms,s,e); ms/=ITERS;
|
||||||
|
|
||||||
|
double flops = 4.0*B*Hq*(double)ql*kl*D;
|
||||||
|
if (causal) flops *= 0.5;
|
||||||
|
double tflops = flops/(ms*1e-3)/1e12;
|
||||||
|
// HBM traffic: Q + O (B*Hq*ql*D each) + K + V (B*Hk*kl*D each), bf16.
|
||||||
|
double bytes = 2.0 * (2.0*nQ + 2.0*nKV);
|
||||||
|
double gbps = bytes/(ms*1e-3)/1e9;
|
||||||
|
|
||||||
|
char cfg[64];
|
||||||
|
snprintf(cfg, sizeof(cfg),
|
||||||
|
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
|
||||||
|
B,Hq,Hk,ql,kl,D,causal);
|
||||||
|
printf("%-46s | %7.4f ms | %7.1f GB/s | %6.2f TFLOP/s\n",
|
||||||
|
cfg, ms, gbps, tflops);
|
||||||
|
|
||||||
|
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
|
||||||
|
delete[]tmp; cudaEventDestroy(s); cudaEventDestroy(e);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
int main() {
|
||||||
|
const int configs[][7] = {
|
||||||
|
{1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal
|
||||||
|
{1,32,4,512,512,128,0}, // standard
|
||||||
|
{1,32,4,128,256,128,0}, // medium
|
||||||
|
{1,4,2,256,256,128,1}, // causal
|
||||||
|
};
|
||||||
|
int n_configs = sizeof(configs) / sizeof(configs[0]);
|
||||||
|
|
||||||
|
for (int ci = 0; ci < n_configs; ci++) {
|
||||||
|
int B=configs[ci][0], Hq=configs[ci][1], Hk=configs[ci][2];
|
||||||
|
int ql=configs[ci][3], kl=configs[ci][4], D=configs[ci][5];
|
||||||
|
int causal=configs[ci][6];
|
||||||
|
printf("=== B=%d Hq=%d Hk=%d q=%d kv=%d D=%d causal=%d ===\n",
|
||||||
|
B,Hq,Hk,ql,kl,D,causal);
|
||||||
|
|
||||||
|
size_t nQ = B*Hq*ql*D, nKV = B*Hk*kl*D;
|
||||||
|
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
||||||
|
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
||||||
|
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
||||||
|
|
||||||
|
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||||
|
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||||
|
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||||
|
tmp=new bf16[max(nQ,nKV)];
|
||||||
|
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
||||||
|
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
||||||
|
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
||||||
|
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
AttentionParams<bf16> p;
|
||||||
|
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
|
||||||
|
p.use_mask=0; p.is_causal=causal; p.causal_offset=0;
|
||||||
|
p.scale=1.0f/sqrtf((float)D);
|
||||||
|
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||||
|
|
||||||
|
double t0=now_ms();
|
||||||
|
dispatch_prefill(p);
|
||||||
|
cudaDeviceSynchronize();
|
||||||
|
double kms=now_ms()-t0;
|
||||||
|
cudaError_t err=cudaGetLastError();
|
||||||
|
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
||||||
|
|
||||||
|
bf16* hOut=new bf16[nQ];
|
||||||
|
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
||||||
|
|
||||||
|
float* ref=new float[nQ];
|
||||||
|
cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal, 0);
|
||||||
|
|
||||||
|
float max_err=0;
|
||||||
|
for (size_t i=0;i<nQ;i++) {
|
||||||
|
float d=fabsf(bf2f(hOut[i])-ref[i]);
|
||||||
|
if(d>max_err) max_err=d;
|
||||||
|
}
|
||||||
|
printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err);
|
||||||
|
|
||||||
|
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
|
||||||
|
delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp;
|
||||||
|
}
|
||||||
|
printf("All tests passed!\n");
|
||||||
|
bench();
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
@@ -0,0 +1,154 @@
|
|||||||
|
#pragma once
|
||||||
|
|
||||||
|
#include <cstdio>
|
||||||
|
#include <cstdlib>
|
||||||
|
#include <cmath>
|
||||||
|
#include <chrono>
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
|
||||||
|
using bf16 = __nv_bfloat16;
|
||||||
|
|
||||||
|
inline bf16 f2bf(float x) { return __float2bfloat16(x); }
|
||||||
|
inline float bf2f(bf16 x) { return __bfloat162float(x); }
|
||||||
|
|
||||||
|
inline float randf() { return (float)rand() / (float)RAND_MAX - 0.5f; }
|
||||||
|
|
||||||
|
inline double now_ms() {
|
||||||
|
using namespace std::chrono;
|
||||||
|
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) \
|
||||||
|
do { \
|
||||||
|
cudaError_t _e = (call); \
|
||||||
|
if (_e != cudaSuccess) { \
|
||||||
|
printf("CUDA error %s at %s:%d\n", cudaGetErrorString(_e), __FILE__, __LINE__); \
|
||||||
|
exit(1); \
|
||||||
|
} \
|
||||||
|
} while (0)
|
||||||
|
|
||||||
|
struct BenchResult {
|
||||||
|
float ms;
|
||||||
|
double gbps;
|
||||||
|
double tflops;
|
||||||
|
};
|
||||||
|
|
||||||
|
template <typename Fn>
|
||||||
|
BenchResult bench_kernel(Fn launch, int warmup, int iters,
|
||||||
|
double flops, double bytes) {
|
||||||
|
for (int i = 0; i < warmup; i++) launch();
|
||||||
|
cudaDeviceSynchronize();
|
||||||
|
cudaError_t err = cudaGetLastError();
|
||||||
|
if (err != cudaSuccess) {
|
||||||
|
printf("CUDA error before bench: %s\n", cudaGetErrorString(err));
|
||||||
|
return {0, 0, 0};
|
||||||
|
}
|
||||||
|
|
||||||
|
cudaEvent_t s, e;
|
||||||
|
cudaEventCreate(&s); cudaEventCreate(&e);
|
||||||
|
cudaEventRecord(s);
|
||||||
|
for (int i = 0; i < iters; i++) launch();
|
||||||
|
cudaEventRecord(e); cudaEventSynchronize(e);
|
||||||
|
float ms = 0; cudaEventElapsedTime(&ms, s, e); ms /= iters;
|
||||||
|
cudaEventDestroy(s); cudaEventDestroy(e);
|
||||||
|
|
||||||
|
return {ms, bytes / (ms * 1e-3) / 1e9, flops / (ms * 1e-3) / 1e12};
|
||||||
|
}
|
||||||
|
|
||||||
|
inline void print_bench_header() {
|
||||||
|
printf("%-46s | %10s | %10s | %10s\n",
|
||||||
|
"config", "latency", "bandwidth", "throughput");
|
||||||
|
printf("---------------------------------------------------------------"
|
||||||
|
"----------------------------\n");
|
||||||
|
}
|
||||||
|
|
||||||
|
inline void print_bench_row(const char* cfg, const BenchResult& r) {
|
||||||
|
printf("%-46s | %7.4f ms | %7.1f GB/s | %6.2f TFLOP/s\n",
|
||||||
|
cfg, r.ms, r.gbps, r.tflops);
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int... Ds>
|
||||||
|
struct _HeadSwitch;
|
||||||
|
|
||||||
|
template <int D>
|
||||||
|
struct _HeadSwitch<D> {
|
||||||
|
template <typename Fn>
|
||||||
|
static void call(int hd, Fn&& fn) { if (hd == D) fn.template operator()<D>(); }
|
||||||
|
};
|
||||||
|
|
||||||
|
template <int D, int... Rest>
|
||||||
|
struct _HeadSwitch<D, Rest...> {
|
||||||
|
template <typename Fn>
|
||||||
|
static void call(int hd, Fn&& fn) {
|
||||||
|
if (hd == D) fn.template operator()<D>();
|
||||||
|
else _HeadSwitch<Rest...>::call(hd, fn);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
// Default set: 32, 64, 128, 256
|
||||||
|
template <typename Fn>
|
||||||
|
void dispatch_by_head_dim(int head_dim, Fn&& fn) {
|
||||||
|
_HeadSwitch<32, 64, 128, 256>::call(head_dim, fn);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Generic CPU reference for multi-query / grouped-query attention.
|
||||||
|
// Tensor shapes (all float*):
|
||||||
|
// Q : [B, Hq, q_len, D]
|
||||||
|
// K : [B, Hk, kv_len, D]
|
||||||
|
// V : [B, Hk, kv_len, D]
|
||||||
|
// O : [B, Hq, q_len, D]
|
||||||
|
// mask: if q_len == 1, shape is [B, kv_len]; otherwise mask is not supported.
|
||||||
|
static void cpu_attention_ref(
|
||||||
|
const float* Q, const float* K, const float* V, const bool* mask,
|
||||||
|
float* O, int B, int Hq, int Hk, int q_len, int kv_len, int D,
|
||||||
|
int is_causal, int causal_offset
|
||||||
|
) {
|
||||||
|
float scale = 1.0f / sqrtf((float)D);
|
||||||
|
int n_rep = Hq / Hk;
|
||||||
|
for (int b = 0; b < B; b++) {
|
||||||
|
for (int h = 0; h < Hq; h++) {
|
||||||
|
int kv_h = h / n_rep;
|
||||||
|
for (int qi = 0; qi < q_len; qi++) {
|
||||||
|
float mv = -INFINITY, sv = 0.0f;
|
||||||
|
float accum[256] = {0.0f};
|
||||||
|
int lim = kv_len;
|
||||||
|
if (is_causal) {
|
||||||
|
int c = qi + causal_offset + 1;
|
||||||
|
lim = (c < kv_len) ? c : kv_len;
|
||||||
|
}
|
||||||
|
for (int kj = 0; kj < lim; kj++) {
|
||||||
|
if (mask != nullptr && q_len == 1) {
|
||||||
|
if (!mask[b * kv_len + kj]) continue;
|
||||||
|
}
|
||||||
|
float dot = 0.0f;
|
||||||
|
size_t q_idx = ((size_t)b * Hq + h) * q_len + qi;
|
||||||
|
size_t kv_idx = ((size_t)b * Hk + kv_h) * kv_len + kj;
|
||||||
|
for (int d = 0; d < D; d++)
|
||||||
|
dot += Q[q_idx * D + d] * K[kv_idx * D + d];
|
||||||
|
dot *= scale;
|
||||||
|
float nm = fmaxf(mv, dot);
|
||||||
|
float a = expf(mv - nm);
|
||||||
|
float b_exp = expf(dot - nm);
|
||||||
|
sv = sv * a + b_exp;
|
||||||
|
for (int d = 0; d < D; d++)
|
||||||
|
accum[d] = accum[d] * a + V[kv_idx * D + d] * b_exp;
|
||||||
|
mv = nm;
|
||||||
|
}
|
||||||
|
float inv = 1.0f / sv;
|
||||||
|
size_t o_idx = ((size_t)b * Hq + h) * q_len + qi;
|
||||||
|
for (int d = 0; d < D; d++)
|
||||||
|
O[o_idx * D + d] = accum[d] * inv;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
+272
-213
@@ -1,22 +1,22 @@
|
|||||||
"""HumanEval code generation benchmark.
|
"""HumanEval benchmark — functional pipeline design.
|
||||||
|
|
||||||
Generates n completions per problem, extracts function bodies, executes
|
Pipeline:
|
||||||
against hidden tests, and computes pass@k.
|
load -> generate -> extract -> test -> score -> report
|
||||||
|
|
||||||
Usage::
|
Each stage is a pure function (except GPU/CPU-bound I/O stages).
|
||||||
|
Config is a single dataclass; side effects are isolated at pipeline boundaries.
|
||||||
python scripts/tools/evaluate_humaneval.py --param_path ./params \
|
|
||||||
--data_path HumanEval.jsonl.gz --output results.json \
|
|
||||||
--num_samples 200 --temperature 0.8 --max_tokens 512
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import argparse
|
import argparse
|
||||||
|
import itertools
|
||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
|
import subprocess
|
||||||
|
import sys
|
||||||
|
from dataclasses import dataclass
|
||||||
from math import prod
|
from math import prod
|
||||||
from multiprocessing import Process, Queue
|
from typing import Dict, Iterator, List, Optional, Sequence, Tuple
|
||||||
from typing import Dict, List, Optional, Tuple
|
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
@@ -26,11 +26,15 @@ from astrai.inference import InferenceEngine
|
|||||||
from astrai.model import AutoModel
|
from astrai.model import AutoModel
|
||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Config
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
HUMANEVAL_URL = (
|
HUMANEVAL_URL = (
|
||||||
"https://github.com/openai/human-eval/raw/master/data/HumanEval.jsonl.gz"
|
"https://github.com/openai/human-eval/raw/master/data/HumanEval.jsonl.gz"
|
||||||
)
|
)
|
||||||
|
|
||||||
_STOP_SEQUENCES = [
|
STOP_SEQUENCES = [
|
||||||
"\nclass ",
|
"\nclass ",
|
||||||
"\ndef ",
|
"\ndef ",
|
||||||
"\n# ",
|
"\n# ",
|
||||||
@@ -40,43 +44,85 @@ _STOP_SEQUENCES = [
|
|||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
def _download_humaneval(data_path: str):
|
@dataclass
|
||||||
if os.path.exists(data_path):
|
class EvalConfig:
|
||||||
|
param_path: str = "./params"
|
||||||
|
data_path: str = "./humaneval/HumanEval.jsonl"
|
||||||
|
output: Optional[str] = None
|
||||||
|
|
||||||
|
test_only: Optional[str] = None
|
||||||
|
generate_only: bool = False
|
||||||
|
|
||||||
|
num_samples: int = 200
|
||||||
|
max_tokens: int = 512
|
||||||
|
temperature: float = 0.8
|
||||||
|
top_p: float = 0.95
|
||||||
|
top_k: int = 50
|
||||||
|
batch_size: int = 32
|
||||||
|
test_timeout: float = 3.0
|
||||||
|
test_workers: int = 8
|
||||||
|
k_values: Tuple[int, ...] = (1, 10, 100)
|
||||||
|
problem_indices: Optional[List[int]] = None
|
||||||
|
|
||||||
|
|
||||||
|
def download(url: str, path: str):
|
||||||
|
if os.path.exists(path):
|
||||||
return
|
return
|
||||||
import gzip
|
import gzip
|
||||||
import urllib.request
|
import urllib.request
|
||||||
|
|
||||||
os.makedirs(os.path.dirname(data_path) or ".", exist_ok=True)
|
os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
|
||||||
print(f"Downloading HumanEval from {HUMANEVAL_URL} ...")
|
print(f"Downloading {url} ...")
|
||||||
tmp = data_path + ".tmp"
|
tmp = path + ".tmp"
|
||||||
urllib.request.urlretrieve(HUMANEVAL_URL, tmp)
|
urllib.request.urlretrieve(url, tmp)
|
||||||
with gzip.open(tmp, "rb") as f_in:
|
with gzip.open(tmp, "rb") as f_in:
|
||||||
with open(data_path, "wb") as f_out:
|
with open(path, "wb") as f_out:
|
||||||
f_out.write(f_in.read())
|
f_out.write(f_in.read())
|
||||||
os.remove(tmp)
|
os.remove(tmp)
|
||||||
print(f" saved to {data_path}")
|
print(f" saved to {path}")
|
||||||
|
|
||||||
|
|
||||||
def _load_problems(data_path: str) -> List[dict]:
|
def load_jsonl(path: str) -> List[dict]:
|
||||||
problems = []
|
rows = []
|
||||||
with open(data_path, "r", encoding="utf-8") as f:
|
with open(path, encoding="utf-8") as f:
|
||||||
for line in f:
|
for line in f:
|
||||||
line = line.strip()
|
line = line.strip()
|
||||||
if line:
|
if line:
|
||||||
problems.append(json.loads(line))
|
rows.append(json.loads(line))
|
||||||
return problems
|
return rows
|
||||||
|
|
||||||
|
|
||||||
def _extract_function_body(code: str, entry_point: str) -> Optional[str]:
|
def save_json(path: str, data):
|
||||||
"""Extract the function body from a completion."""
|
with open(path, "w", encoding="utf-8") as f:
|
||||||
|
json.dump(data, f, indent=2, ensure_ascii=False)
|
||||||
|
|
||||||
|
|
||||||
|
def create_engine(param_path: str, batch_size: int) -> InferenceEngine:
|
||||||
|
model = AutoModel.from_pretrained(param_path)
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
||||||
|
model.to(device="cuda", dtype=torch.bfloat16)
|
||||||
|
return InferenceEngine(
|
||||||
|
model=model,
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
max_batch_size=batch_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def trim_stop(text: str) -> str:
|
||||||
|
for stop in STOP_SEQUENCES:
|
||||||
|
idx = text.find(stop)
|
||||||
|
if idx != -1:
|
||||||
|
text = text[:idx]
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
def extract_body(code: str, entry_point: str) -> Optional[str]:
|
||||||
pattern = rf"def\s+{re.escape(entry_point)}\b[^:]*:"
|
pattern = rf"def\s+{re.escape(entry_point)}\b[^:]*:"
|
||||||
match = re.search(pattern, code)
|
match = re.search(pattern, code)
|
||||||
if not match:
|
if not match:
|
||||||
# Use the full code as-is if we can't find the function
|
|
||||||
return code
|
return code
|
||||||
|
|
||||||
body_start = match.end()
|
lines = code[match.end() :].split("\n")
|
||||||
lines = code[body_start:].split("\n")
|
|
||||||
body_lines = []
|
body_lines = []
|
||||||
started = False
|
started = False
|
||||||
|
|
||||||
@@ -94,240 +140,253 @@ def _extract_function_body(code: str, entry_point: str) -> Optional[str]:
|
|||||||
body_lines.append(stripped)
|
body_lines.append(stripped)
|
||||||
|
|
||||||
body = "\n".join(body_lines)
|
body = "\n".join(body_lines)
|
||||||
if not body.strip():
|
return body if body.strip() else None
|
||||||
return None
|
|
||||||
return body
|
|
||||||
|
|
||||||
|
|
||||||
def _trim_stop_sequences(text: str) -> str:
|
def deduplicate(seq: Sequence[str]) -> List[str]:
|
||||||
for stop in _STOP_SEQUENCES:
|
|
||||||
idx = text.find(stop)
|
|
||||||
if idx != -1:
|
|
||||||
text = text[:idx]
|
|
||||||
return text
|
|
||||||
|
|
||||||
|
|
||||||
def _execute_code(problem: dict, completion: str, timeout: float = 3.0) -> bool:
|
|
||||||
"""Run the completion against hidden tests in a subprocess."""
|
|
||||||
|
|
||||||
def _worker(queue, full_code):
|
|
||||||
try:
|
|
||||||
namespace = {}
|
|
||||||
exec(full_code, namespace)
|
|
||||||
check = namespace.get("check")
|
|
||||||
if check is None:
|
|
||||||
queue.put(False)
|
|
||||||
return
|
|
||||||
check(namespace.get(problem["entry_point"]))
|
|
||||||
queue.put(True)
|
|
||||||
except Exception:
|
|
||||||
queue.put(False)
|
|
||||||
|
|
||||||
full_code = problem["prompt"] + completion + "\n" + problem["test"]
|
|
||||||
|
|
||||||
queue: Queue = Queue()
|
|
||||||
proc = Process(target=_worker, args=(queue, full_code))
|
|
||||||
proc.start()
|
|
||||||
proc.join(timeout)
|
|
||||||
|
|
||||||
if proc.is_alive():
|
|
||||||
proc.terminate()
|
|
||||||
proc.join()
|
|
||||||
return False
|
|
||||||
|
|
||||||
try:
|
|
||||||
return queue.get_nowait()
|
|
||||||
except Exception:
|
|
||||||
return False
|
|
||||||
|
|
||||||
|
|
||||||
def _pass_at_k(n: int, c: int, k: int) -> float:
|
|
||||||
"""Unbiased estimator of pass@k."""
|
|
||||||
if n - c < k:
|
|
||||||
return 1.0
|
|
||||||
return 1.0 - float(prod(1.0 - k / np.arange(n - c + 1, n + 1)))
|
|
||||||
|
|
||||||
|
|
||||||
def _deduplicate(completions: List[str]) -> List[str]:
|
|
||||||
seen = set()
|
seen = set()
|
||||||
unique = []
|
return [x for x in seq if not (x in seen or seen.add(x))]
|
||||||
for c in completions:
|
|
||||||
if c not in seen:
|
|
||||||
seen.add(c)
|
|
||||||
unique.append(c)
|
|
||||||
return unique
|
|
||||||
|
|
||||||
|
|
||||||
def _generate(
|
def generate_batch(
|
||||||
engine: InferenceEngine,
|
engine: InferenceEngine,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
num_samples: int,
|
n: int,
|
||||||
|
batch_size: int,
|
||||||
max_tokens: int,
|
max_tokens: int,
|
||||||
temperature: float,
|
temperature: float,
|
||||||
top_p: float,
|
top_p: float,
|
||||||
top_k: int,
|
top_k: int,
|
||||||
batch_size: int,
|
|
||||||
) -> List[str]:
|
) -> List[str]:
|
||||||
batches = [prompt] * min(batch_size, num_samples)
|
|
||||||
completions = []
|
completions = []
|
||||||
remaining = num_samples
|
remaining = n
|
||||||
|
|
||||||
while remaining > 0:
|
while remaining > 0:
|
||||||
current = min(batch_size, remaining)
|
current = min(batch_size, remaining)
|
||||||
batch_prompts = batches[:current]
|
|
||||||
outputs = engine.generate(
|
outputs = engine.generate(
|
||||||
prompt=batch_prompts,
|
prompt=[prompt] * current,
|
||||||
stream=False,
|
stream=False,
|
||||||
max_tokens=max_tokens,
|
max_tokens=max_tokens,
|
||||||
temperature=temperature,
|
temperature=temperature,
|
||||||
top_p=top_p,
|
top_p=top_p,
|
||||||
top_k=top_k,
|
top_k=top_k,
|
||||||
)
|
)
|
||||||
if isinstance(outputs, str):
|
completions.extend(outputs if isinstance(outputs, list) else [outputs])
|
||||||
outputs = [outputs]
|
|
||||||
completions.extend(outputs)
|
|
||||||
remaining -= current
|
remaining -= current
|
||||||
|
return deduplicate(completions)
|
||||||
return _deduplicate(completions)
|
|
||||||
|
|
||||||
|
|
||||||
def evaluate(
|
def extract_completions(
|
||||||
|
raw: Sequence[str],
|
||||||
|
entry_point: str,
|
||||||
|
) -> List[str]:
|
||||||
|
bodies = []
|
||||||
|
for r in raw:
|
||||||
|
t = trim_stop(r)
|
||||||
|
body = extract_body(t, entry_point)
|
||||||
|
if body:
|
||||||
|
bodies.append(body)
|
||||||
|
return bodies
|
||||||
|
|
||||||
|
|
||||||
|
def generate_all(
|
||||||
engine: InferenceEngine,
|
engine: InferenceEngine,
|
||||||
problems: List[dict],
|
problems: Sequence[dict],
|
||||||
num_samples: int,
|
cfg: EvalConfig,
|
||||||
max_tokens: int,
|
) -> List[dict]:
|
||||||
temperature: float,
|
results = []
|
||||||
top_p: float,
|
for problem in tqdm.tqdm(problems, desc="Generating", unit="problem"):
|
||||||
top_k: int,
|
raw = generate_batch(
|
||||||
batch_size: int,
|
|
||||||
k_values: Tuple[int, ...] = (1, 10, 100),
|
|
||||||
) -> Dict:
|
|
||||||
results = {}
|
|
||||||
all_pass_at_k = {k: [] for k in k_values}
|
|
||||||
|
|
||||||
for problem in tqdm.tqdm(problems, desc="HumanEval", unit="problem"):
|
|
||||||
task_id = problem["task_id"]
|
|
||||||
prompt = problem["prompt"]
|
|
||||||
entry_point = problem["entry_point"]
|
|
||||||
|
|
||||||
raw_completions = _generate(
|
|
||||||
engine,
|
engine,
|
||||||
prompt,
|
problem["prompt"],
|
||||||
num_samples,
|
cfg.num_samples,
|
||||||
max_tokens,
|
cfg.batch_size,
|
||||||
temperature,
|
cfg.max_tokens,
|
||||||
top_p,
|
cfg.temperature,
|
||||||
top_k,
|
cfg.top_p,
|
||||||
batch_size,
|
cfg.top_k,
|
||||||
|
)
|
||||||
|
bodies = extract_completions(raw, problem["entry_point"])
|
||||||
|
results.append(
|
||||||
|
dict(
|
||||||
|
task_id=problem["task_id"],
|
||||||
|
entry_point=problem["entry_point"],
|
||||||
|
prompt=problem["prompt"],
|
||||||
|
test=problem["test"],
|
||||||
|
completions=bodies,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
completions = []
|
|
||||||
for raw in raw_completions:
|
|
||||||
trimmed = _trim_stop_sequences(raw)
|
|
||||||
body = _extract_function_body(trimmed, entry_point)
|
|
||||||
if body:
|
|
||||||
completions.append(body)
|
|
||||||
|
|
||||||
passed = 0
|
|
||||||
for comp in completions:
|
|
||||||
if _execute_code(problem, comp):
|
|
||||||
passed += 1
|
|
||||||
|
|
||||||
n = len(completions)
|
|
||||||
c = passed
|
|
||||||
result = {"task_id": task_id, "n": n, "passed": c}
|
|
||||||
for k in k_values:
|
|
||||||
result[f"pass@{k}"] = round(_pass_at_k(n, c, k), 4)
|
|
||||||
all_pass_at_k[k].append(_pass_at_k(n, c, k))
|
|
||||||
results[task_id] = result
|
|
||||||
|
|
||||||
summary = {}
|
|
||||||
for k in k_values:
|
|
||||||
vals = all_pass_at_k[k]
|
|
||||||
summary[f"pass@{k}"] = round(float(np.mean(vals)), 4)
|
|
||||||
results["_summary"] = summary
|
|
||||||
|
|
||||||
return results
|
return results
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def execute_one(args: tuple) -> bool:
|
||||||
parser = argparse.ArgumentParser(description="HumanEval benchmark")
|
full_code, entry_point, timeout = args
|
||||||
parser.add_argument(
|
try:
|
||||||
"--param_path", type=str, default="./params", help="Model directory"
|
r = subprocess.run(
|
||||||
)
|
[sys.executable, "-c", full_code],
|
||||||
parser.add_argument(
|
capture_output=True,
|
||||||
"--data_path",
|
timeout=timeout,
|
||||||
|
)
|
||||||
|
return r.returncode == 0
|
||||||
|
except subprocess.TimeoutExpired:
|
||||||
|
return False
|
||||||
|
except Exception:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def test_one(item: dict, cfg: EvalConfig) -> Tuple[str, int, int]:
|
||||||
|
from concurrent.futures import ProcessPoolExecutor
|
||||||
|
|
||||||
|
task_id = item["task_id"]
|
||||||
|
completions = item["completions"]
|
||||||
|
codes = [
|
||||||
|
(
|
||||||
|
item["prompt"] + c + "\n" + item["test"],
|
||||||
|
item["entry_point"],
|
||||||
|
cfg.test_timeout,
|
||||||
|
)
|
||||||
|
for c in completions
|
||||||
|
]
|
||||||
|
n = len(codes)
|
||||||
|
passed = 0
|
||||||
|
with ProcessPoolExecutor(max_workers=cfg.test_workers) as pool:
|
||||||
|
for ok in pool.map(execute_one, codes):
|
||||||
|
if ok:
|
||||||
|
passed += 1
|
||||||
|
return task_id, n, passed
|
||||||
|
|
||||||
|
|
||||||
|
def test_all(
|
||||||
|
items: Sequence[dict],
|
||||||
|
cfg: EvalConfig,
|
||||||
|
) -> Iterator[Tuple[str, int, int]]:
|
||||||
|
for item in tqdm.tqdm(items, desc="Testing", unit="problem"):
|
||||||
|
yield test_one(item, cfg)
|
||||||
|
|
||||||
|
|
||||||
|
def pass_at_k(n: int, c: int, k: int) -> float:
|
||||||
|
if n - c < k:
|
||||||
|
return 1.0
|
||||||
|
return 1.0 - float(prod(1.0 - k / np.arange(n - c + 1, n + 1)))
|
||||||
|
|
||||||
|
|
||||||
|
def score_results(
|
||||||
|
results: Iterator[Tuple[str, int, int]],
|
||||||
|
k_values: Tuple[int, ...],
|
||||||
|
) -> Dict:
|
||||||
|
# filter to k <= n (peek first result to get n)
|
||||||
|
first = next(results)
|
||||||
|
results = itertools.chain([first], results)
|
||||||
|
n = first[1]
|
||||||
|
k_values = tuple(k for k in k_values if k <= n)
|
||||||
|
|
||||||
|
scores = {k: [] for k in k_values}
|
||||||
|
output = {}
|
||||||
|
for task_id, n, passed in results:
|
||||||
|
entry = {"task_id": task_id, "n": n, "passed": passed}
|
||||||
|
for k in k_values:
|
||||||
|
pk = round(pass_at_k(n, passed, k), 4)
|
||||||
|
entry[f"pass@{k}"] = pk
|
||||||
|
scores[k].append(pk)
|
||||||
|
output[task_id] = entry
|
||||||
|
|
||||||
|
summary = {}
|
||||||
|
for k in k_values:
|
||||||
|
vals = scores[k]
|
||||||
|
summary[f"pass@{k}"] = round(float(np.mean(vals)), 4)
|
||||||
|
output["_summary"] = summary
|
||||||
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
def run_pipeline(cfg: EvalConfig) -> Dict:
|
||||||
|
if cfg.test_only:
|
||||||
|
with open(cfg.test_only, encoding="utf-8") as f:
|
||||||
|
generated = json.load(f)
|
||||||
|
else:
|
||||||
|
download(HUMANEVAL_URL, cfg.data_path)
|
||||||
|
|
||||||
|
problems = load_jsonl(cfg.data_path)
|
||||||
|
if cfg.problem_indices:
|
||||||
|
problems = [problems[i] for i in cfg.problem_indices if i < len(problems)]
|
||||||
|
|
||||||
|
engine = create_engine(cfg.param_path, cfg.batch_size)
|
||||||
|
|
||||||
|
try:
|
||||||
|
generated = generate_all(engine, problems, cfg)
|
||||||
|
finally:
|
||||||
|
engine.shutdown()
|
||||||
|
|
||||||
|
if cfg.output:
|
||||||
|
mid = cfg.output.replace(".json", "_completions.json")
|
||||||
|
save_json(mid, generated)
|
||||||
|
print(f"Completions saved to {mid}")
|
||||||
|
|
||||||
|
if cfg.generate_only:
|
||||||
|
return {}
|
||||||
|
|
||||||
|
results = test_all(generated, cfg)
|
||||||
|
scored = score_results(results, cfg.k_values)
|
||||||
|
return scored
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args(argv: Optional[List[str]] = None) -> EvalConfig:
|
||||||
|
p = argparse.ArgumentParser(description="HumanEval benchmark")
|
||||||
|
p.add_argument("--param_path", type=str, default="./params")
|
||||||
|
p.add_argument("--data_path", type=str, default="./humaneval/HumanEval.jsonl")
|
||||||
|
p.add_argument("--output", type=str, default=None)
|
||||||
|
p.add_argument(
|
||||||
|
"--test_only",
|
||||||
type=str,
|
type=str,
|
||||||
default="./humaneval/HumanEval.jsonl",
|
|
||||||
help="HumanEval JSONL file (auto-download if missing)",
|
|
||||||
)
|
|
||||||
parser.add_argument("--output", type=str, default=None, help="Output JSON path")
|
|
||||||
parser.add_argument(
|
|
||||||
"--num_samples",
|
|
||||||
type=int,
|
|
||||||
default=200,
|
|
||||||
help="Completions per problem",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--max_tokens", type=int, default=512, help="Max generation tokens"
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--temperature", type=float, default=0.8, help="Sampling temperature"
|
|
||||||
)
|
|
||||||
parser.add_argument("--top_p", type=float, default=0.95, help="Top-p sampling")
|
|
||||||
parser.add_argument("--top_k", type=int, default=50, help="Top-k sampling")
|
|
||||||
parser.add_argument(
|
|
||||||
"--batch_size", type=int, default=1, help="Inference batch size"
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--problems",
|
|
||||||
type=int,
|
|
||||||
nargs="+",
|
|
||||||
default=None,
|
default=None,
|
||||||
help="Specific problem indices (0-based)",
|
help="Skip generation, test existing completions JSON",
|
||||||
)
|
)
|
||||||
args = parser.parse_args()
|
p.add_argument(
|
||||||
|
"--generate_only", action="store_true", help="Only generate, skip testing"
|
||||||
_download_humaneval(args.data_path)
|
|
||||||
problems = _load_problems(args.data_path)
|
|
||||||
if args.problems:
|
|
||||||
problems = [problems[i] for i in args.problems if i < len(problems)]
|
|
||||||
|
|
||||||
model = AutoModel.from_pretrained(args.param_path)
|
|
||||||
tokenizer = AutoTokenizer.from_pretrained(args.param_path)
|
|
||||||
model.to(device="cuda", dtype=torch.bfloat16)
|
|
||||||
|
|
||||||
engine = InferenceEngine(
|
|
||||||
model=model,
|
|
||||||
tokenizer=tokenizer,
|
|
||||||
max_batch_size=args.batch_size,
|
|
||||||
)
|
)
|
||||||
|
p.add_argument("--num_samples", type=int, default=200)
|
||||||
|
p.add_argument("--max_tokens", type=int, default=512)
|
||||||
|
p.add_argument("--temperature", type=float, default=0.8)
|
||||||
|
p.add_argument("--top_p", type=float, default=0.95)
|
||||||
|
p.add_argument("--top_k", type=int, default=50)
|
||||||
|
p.add_argument("--batch_size", type=int, default=32)
|
||||||
|
p.add_argument("--test_workers", type=int, default=8)
|
||||||
|
p.add_argument("--test_timeout", type=float, default=3.0)
|
||||||
|
p.add_argument("--problems", type=int, nargs="+", default=None)
|
||||||
|
args = p.parse_args(argv)
|
||||||
|
|
||||||
results = evaluate(
|
return EvalConfig(
|
||||||
engine=engine,
|
param_path=args.param_path,
|
||||||
problems=problems,
|
data_path=args.data_path,
|
||||||
|
output=args.output,
|
||||||
|
test_only=args.test_only,
|
||||||
|
generate_only=args.generate_only,
|
||||||
num_samples=args.num_samples,
|
num_samples=args.num_samples,
|
||||||
max_tokens=args.max_tokens,
|
max_tokens=args.max_tokens,
|
||||||
temperature=args.temperature,
|
temperature=args.temperature,
|
||||||
top_p=args.top_p,
|
top_p=args.top_p,
|
||||||
top_k=args.top_k,
|
top_k=args.top_k,
|
||||||
batch_size=args.batch_size,
|
batch_size=args.batch_size,
|
||||||
k_values=(1, 10, 100),
|
test_workers=args.test_workers,
|
||||||
|
test_timeout=args.test_timeout,
|
||||||
|
problem_indices=args.problems,
|
||||||
)
|
)
|
||||||
|
|
||||||
summary = results.pop("_summary")
|
|
||||||
|
def report(scored: Dict):
|
||||||
|
summary = scored.pop("_summary", {})
|
||||||
print(f"\n{'=' * 60}")
|
print(f"\n{'=' * 60}")
|
||||||
for k, v in summary.items():
|
for k, v in summary.items():
|
||||||
print(f" {k}: {v:.2%}")
|
print(f" {k}: {v:.2%}")
|
||||||
print(f"{'=' * 60}")
|
print(f"{'=' * 60}")
|
||||||
|
scored["_summary"] = summary
|
||||||
|
|
||||||
if args.output:
|
|
||||||
results["_summary"] = summary
|
|
||||||
with open(args.output, "w", encoding="utf-8") as f:
|
|
||||||
json.dump(results, f, indent=2, ensure_ascii=False)
|
|
||||||
print(f"Results saved to {args.output}")
|
|
||||||
|
|
||||||
engine.shutdown()
|
def main():
|
||||||
|
cfg = parse_args()
|
||||||
|
scored = run_pipeline(cfg)
|
||||||
|
report(scored)
|
||||||
|
if cfg.output:
|
||||||
|
save_json(cfg.output, scored)
|
||||||
|
print(f"Results saved to {cfg.output}")
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
+384
-217
@@ -1,31 +1,23 @@
|
|||||||
"""IFD (Instruction Following Difficulty) data quality scoring.
|
"""IFD (Instruction Following Difficulty) data quality scoring.
|
||||||
|
|
||||||
Computes IFD scores for instruction-response pairs to guide data selection.
|
IFD = conditional_NLL / unconditional_NLL
|
||||||
IFD = conditional_NLL / unconditional_NLL, where:
|
|
||||||
|
|
||||||
- conditional_NLL: average CE loss on response tokens given instruction context
|
- Messages format: plain text concatenation (no chat template)
|
||||||
- unconditional_NLL: average CE loss on response tokens alone
|
- Plain format: raw instr_key + resp_key fields
|
||||||
|
|
||||||
Higher IFD (close to 1) = instruction provides less help = harder sample.
|
v2 changelog:
|
||||||
Lower IFD (close to 0) = instruction provides strong guidance = easy sample.
|
- Same token set: unconditional pass prefixes resp with a plain-text sentinel
|
||||||
IFD > 1 = instruction misleads the model = likely low-quality data.
|
(default ``\\n``; use ``--sentinel_text ""`` for bos/pad fallback).
|
||||||
|
Both branches predict the identical N resp tokens.
|
||||||
Usage::
|
Single-token answers (rl=1) are now supported.
|
||||||
|
- ctx_len tracked in output
|
||||||
python scripts/eval/ifd.py --param_path ./params \
|
- skip_reason for None samples (no more silent None)
|
||||||
--input data.jsonl --output data_with_ifd.jsonl \
|
- --per_token for per-token IFD breakdown
|
||||||
--instr_key instruction --resp_key response
|
|
||||||
|
|
||||||
Disable chat template::
|
|
||||||
|
|
||||||
python scripts/eval/ifd.py --param_path ./params \
|
|
||||||
--input data.jsonl --output data_with_ifd.jsonl \
|
|
||||||
--instr_key instruction --resp_key response \
|
|
||||||
--no_chat_template
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import argparse
|
import argparse
|
||||||
import json
|
import json
|
||||||
|
import statistics
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
@@ -35,230 +27,396 @@ from astrai.model import AutoModel
|
|||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
|
||||||
def compute_ifd(
|
def _pack_bins(pairs, max_len):
|
||||||
|
"""BFD bin packing: pack (c+r) into bins of max total length."""
|
||||||
|
indexed = sorted(enumerate(pairs), key=lambda x: -(len(x[1][0]) + len(x[1][1])))
|
||||||
|
bins = []
|
||||||
|
lengths = []
|
||||||
|
for orig_idx, (c, r) in indexed:
|
||||||
|
size = len(c) + len(r)
|
||||||
|
best_bin = -1
|
||||||
|
for bi, rem in enumerate(lengths):
|
||||||
|
if rem >= size:
|
||||||
|
if best_bin < 0 or rem < lengths[best_bin]:
|
||||||
|
best_bin = bi
|
||||||
|
if best_bin >= 0:
|
||||||
|
bins[best_bin].append((orig_idx, c, r))
|
||||||
|
lengths[best_bin] -= size
|
||||||
|
else:
|
||||||
|
bins.append([(orig_idx, c, r)])
|
||||||
|
lengths.append(max_len - size)
|
||||||
|
return bins
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_sentinel_ids(tokenizer, sentinel_text):
|
||||||
|
"""Tokenize the sentinel text for the unconditional pass prefix.
|
||||||
|
|
||||||
|
Falls back to bos/pad_token_id when sentinel_text is empty or
|
||||||
|
cannot be encoded.
|
||||||
|
"""
|
||||||
|
if sentinel_text:
|
||||||
|
ids = tokenizer.encode(sentinel_text, add_special_tokens=False)
|
||||||
|
if ids:
|
||||||
|
return ids
|
||||||
|
for attr in ("bos_token_id", "pad_token_id", "eos_token_id"):
|
||||||
|
tid = getattr(tokenizer, attr, None)
|
||||||
|
if tid is not None:
|
||||||
|
return [tid]
|
||||||
|
return [0]
|
||||||
|
|
||||||
|
|
||||||
|
@torch.inference_mode()
|
||||||
|
def _score_batch(
|
||||||
|
pairs, model, device, max_len=2048, sentinel_ids=None, per_token=False
|
||||||
|
):
|
||||||
|
"""BFD-packed IFD with text-sentinel-anchored unconditional pass.
|
||||||
|
|
||||||
|
Conditional: (ctx + resp[0..i-1]) → resp[i], i = 0..N-1
|
||||||
|
Unconditional: (<sentinel> + resp[0..i-1]) → resp[i], i = 0..N-1
|
||||||
|
|
||||||
|
Both branches predict the identical N response tokens. A short
|
||||||
|
plain-text sentinel gives the unconditional pass a prefix so that
|
||||||
|
every response token can be predicted. Single-token answers (rl=1)
|
||||||
|
are supported.
|
||||||
|
"""
|
||||||
|
if not pairs:
|
||||||
|
return []
|
||||||
|
|
||||||
|
if sentinel_ids is None:
|
||||||
|
sentinel_ids = [0]
|
||||||
|
|
||||||
|
bins = _pack_bins(pairs, max_len)
|
||||||
|
result = [None] * len(pairs)
|
||||||
|
|
||||||
|
# ---- conditional pass (packed, per-document position IDs) ----
|
||||||
|
for bin_items in bins:
|
||||||
|
seq_ids = []
|
||||||
|
global_pos = []
|
||||||
|
doc_ids = []
|
||||||
|
doc_offsets = []
|
||||||
|
|
||||||
|
for di, (orig_idx, c, r) in enumerate(bin_items):
|
||||||
|
ctx_len = len(c)
|
||||||
|
start = len(seq_ids)
|
||||||
|
item_len = len(c) + len(r)
|
||||||
|
seq_ids.extend(c)
|
||||||
|
seq_ids.extend(r)
|
||||||
|
end = len(seq_ids)
|
||||||
|
global_pos.extend(range(item_len))
|
||||||
|
doc_ids.extend([di] * item_len)
|
||||||
|
doc_offsets.append((start, end, orig_idx, ctx_len))
|
||||||
|
|
||||||
|
full_ids = torch.tensor([seq_ids], device=device, dtype=torch.long)
|
||||||
|
pos_ids = torch.tensor([global_pos], device=device, dtype=torch.long)
|
||||||
|
seq_len = len(seq_ids)
|
||||||
|
causal = torch.tril(
|
||||||
|
torch.ones(seq_len, seq_len, dtype=torch.bool, device=device)
|
||||||
|
)
|
||||||
|
doc_t = torch.tensor([doc_ids], device=device)
|
||||||
|
doc_mask = doc_t.unsqueeze(-1) == doc_t.unsqueeze(-2)
|
||||||
|
attn_mask = (causal & doc_mask[0]).unsqueeze(0).unsqueeze(0)
|
||||||
|
logits_full = model(full_ids, position_ids=pos_ids, input_mask=attn_mask)[
|
||||||
|
"logits"
|
||||||
|
][0]
|
||||||
|
|
||||||
|
for start, end, orig_idx, ctx_len in doc_offsets:
|
||||||
|
rl = end - start - ctx_len
|
||||||
|
resp_start = start + ctx_len - 1
|
||||||
|
resp_logits = logits_full[resp_start : end - 1]
|
||||||
|
resp_targets = torch.tensor(
|
||||||
|
seq_ids[start + ctx_len : end], device=device, dtype=torch.long
|
||||||
|
)
|
||||||
|
cond_losses = F.cross_entropy(
|
||||||
|
resp_logits, resp_targets, reduction="none"
|
||||||
|
).cpu()
|
||||||
|
result[orig_idx] = {
|
||||||
|
"_cond_losses": cond_losses,
|
||||||
|
"_rl": rl,
|
||||||
|
"_ctx_len": ctx_len,
|
||||||
|
}
|
||||||
|
|
||||||
|
# ---- unconditional pass (sentinel-prefixed, batched 2D) ----
|
||||||
|
valid_items = [
|
||||||
|
(
|
||||||
|
i,
|
||||||
|
result[i]["_rl"],
|
||||||
|
result[i]["_ctx_len"],
|
||||||
|
result[i]["_cond_losses"],
|
||||||
|
pairs[i][1],
|
||||||
|
)
|
||||||
|
for i in range(len(pairs))
|
||||||
|
if result[i] is not None and "_cond_losses" in result[i]
|
||||||
|
]
|
||||||
|
if not valid_items:
|
||||||
|
return result
|
||||||
|
|
||||||
|
valid_items.sort(key=lambda x: -x[1])
|
||||||
|
prefix_len = len(sentinel_ids)
|
||||||
|
max_rl = prefix_len + max(rl for _, rl, _, _, _ in valid_items)
|
||||||
|
bsz = len(valid_items)
|
||||||
|
|
||||||
|
u_batch = torch.zeros(bsz, max_rl, dtype=torch.long, device=device)
|
||||||
|
for ri, (_, rl, _, _, r_ids) in enumerate(valid_items):
|
||||||
|
u_batch[ri, :prefix_len] = torch.tensor(sentinel_ids, dtype=torch.long)
|
||||||
|
u_batch[ri, prefix_len : prefix_len + rl] = torch.tensor(
|
||||||
|
r_ids, dtype=torch.long
|
||||||
|
)
|
||||||
|
|
||||||
|
logits_resp = model(u_batch)["logits"]
|
||||||
|
|
||||||
|
for ri, (orig_idx, rl, ctx_len, cond_losses, _) in enumerate(valid_items):
|
||||||
|
unp_logits = logits_resp[ri, prefix_len - 1 : prefix_len - 1 + rl]
|
||||||
|
unp_targets = u_batch[ri, prefix_len : prefix_len + rl]
|
||||||
|
uncond_losses = F.cross_entropy(unp_logits, unp_targets, reduction="none").cpu()
|
||||||
|
|
||||||
|
L_cond = cond_losses.mean().item()
|
||||||
|
L_uncond = uncond_losses.mean().item()
|
||||||
|
ifd = L_cond / L_uncond if L_uncond > 0 else None
|
||||||
|
|
||||||
|
out = {
|
||||||
|
"L_cond": round(L_cond, 6),
|
||||||
|
"L_uncond": round(L_uncond, 6),
|
||||||
|
"ifd": round(ifd, 6) if ifd is not None else None,
|
||||||
|
"ctx_len": ctx_len,
|
||||||
|
"resp_len": rl,
|
||||||
|
}
|
||||||
|
if per_token:
|
||||||
|
per = [
|
||||||
|
(round(c.item() / u.item(), 6) if u.item() > 0 else None)
|
||||||
|
for c, u in zip(cond_losses, uncond_losses)
|
||||||
|
]
|
||||||
|
out["ifd_per_token"] = per
|
||||||
|
result[orig_idx] = out
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def _trim(context_ids, resp_ids, max_len):
|
||||||
|
"""Truncate to fit max_len, keeping response intact if possible."""
|
||||||
|
if len(resp_ids) > max_len // 2:
|
||||||
|
resp_ids = resp_ids[: max_len // 2]
|
||||||
|
full_ids = context_ids + resp_ids
|
||||||
|
if len(full_ids) <= max_len:
|
||||||
|
return context_ids, resp_ids
|
||||||
|
overflow = len(full_ids) - max_len
|
||||||
|
if overflow >= len(context_ids):
|
||||||
|
return [], resp_ids[:max_len]
|
||||||
|
return context_ids[overflow:], resp_ids
|
||||||
|
|
||||||
|
|
||||||
|
def score_plain(
|
||||||
model,
|
model,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
instruction: str,
|
instruction,
|
||||||
response: str,
|
response,
|
||||||
device: str,
|
device,
|
||||||
max_len: int = 2048,
|
max_len=2048,
|
||||||
use_chat_template: bool = False,
|
sentinel_ids=None,
|
||||||
) -> dict:
|
per_token=False,
|
||||||
if use_chat_template:
|
):
|
||||||
return _compute_ifd_with_template(
|
"""Compute IFD for a single instruction-response pair (plain format)."""
|
||||||
model, tokenizer, instruction, response, device, max_len
|
ctx_ids = tokenizer.encode(instruction, add_special_tokens=False)
|
||||||
)
|
|
||||||
return _compute_ifd_raw(model, tokenizer, instruction, response, device, max_len)
|
|
||||||
|
|
||||||
|
|
||||||
def _compute_ifd_raw(model, tokenizer, instruction, response, device, max_len) -> dict:
|
|
||||||
instr_ids = tokenizer.encode(instruction, add_special_tokens=False)
|
|
||||||
resp_ids = tokenizer.encode(response, add_special_tokens=False)
|
resp_ids = tokenizer.encode(response, add_special_tokens=False)
|
||||||
|
ctx_ids, resp_ids = _trim(ctx_ids, resp_ids, max_len)
|
||||||
if len(resp_ids) > max_len:
|
if not ctx_ids or not resp_ids:
|
||||||
resp_ids = resp_ids[:max_len]
|
|
||||||
|
|
||||||
if not resp_ids:
|
|
||||||
return {
|
return {
|
||||||
"L_cond": None,
|
"L_cond": None,
|
||||||
"L_uncond": None,
|
"L_uncond": None,
|
||||||
"ifd": None,
|
"ifd": None,
|
||||||
"error": "empty response",
|
"skip_reason": "empty ctx or resp",
|
||||||
}
|
}
|
||||||
|
return _score_batch(
|
||||||
qa_len = len(instr_ids) + len(resp_ids)
|
[(ctx_ids, resp_ids)],
|
||||||
if qa_len > max_len:
|
model,
|
||||||
overflow = qa_len - max_len
|
device,
|
||||||
if overflow >= len(instr_ids):
|
max_len,
|
||||||
resp_ids = resp_ids[:max_len]
|
sentinel_ids=sentinel_ids,
|
||||||
instr_ids = []
|
per_token=per_token,
|
||||||
else:
|
)[0]
|
||||||
instr_ids = instr_ids[overflow:]
|
|
||||||
|
|
||||||
if not instr_ids:
|
|
||||||
return {
|
|
||||||
"L_cond": None,
|
|
||||||
"L_uncond": None,
|
|
||||||
"ifd": None,
|
|
||||||
"error": "response too long for context",
|
|
||||||
}
|
|
||||||
|
|
||||||
instr_len = len(instr_ids)
|
|
||||||
resp_len = len(resp_ids)
|
|
||||||
|
|
||||||
qa_ids = instr_ids + resp_ids
|
|
||||||
|
|
||||||
with torch.inference_mode():
|
|
||||||
logits_qa = model(torch.tensor([qa_ids], device=device, dtype=torch.long))[
|
|
||||||
"logits"
|
|
||||||
][0]
|
|
||||||
logits_resp = model(torch.tensor([resp_ids], device=device, dtype=torch.long))[
|
|
||||||
"logits"
|
|
||||||
][0]
|
|
||||||
|
|
||||||
resp_logits = logits_qa[instr_len - 1 : -1]
|
|
||||||
resp_targets = logits_resp.new_tensor(resp_ids, dtype=torch.long)
|
|
||||||
L_cond = F.cross_entropy(resp_logits, resp_targets, reduction="mean").item()
|
|
||||||
|
|
||||||
unp_logits = logits_resp[:-1]
|
|
||||||
unp_targets = logits_resp.new_tensor(resp_ids[1:], dtype=torch.long)
|
|
||||||
L_uncond = F.cross_entropy(unp_logits, unp_targets, reduction="mean").item()
|
|
||||||
|
|
||||||
ifd = L_cond / L_uncond if L_uncond > 0 else None
|
|
||||||
|
|
||||||
return {
|
|
||||||
"L_cond": round(L_cond, 6),
|
|
||||||
"L_uncond": round(L_uncond, 6),
|
|
||||||
"ifd": round(ifd, 6) if ifd is not None else None,
|
|
||||||
"instr_len": instr_len,
|
|
||||||
"resp_len": resp_len,
|
|
||||||
"error": None,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def _compute_ifd_with_template(
|
def score_messages(
|
||||||
model, tokenizer, instruction, response, device, max_len
|
model, tokenizer, messages, device, max_len=2048, sentinel_ids=None, per_token=False
|
||||||
) -> dict:
|
):
|
||||||
instr_prefix = tokenizer.apply_chat_template(
|
"""Compute IFD for each assistant turn in a messages array."""
|
||||||
[{"role": "user", "content": instruction}],
|
turns = []
|
||||||
tokenize=False,
|
for i, msg in enumerate(messages):
|
||||||
add_generation_prompt=True,
|
if msg.get("role") != "assistant":
|
||||||
|
continue
|
||||||
|
ctx_text = "\n\n".join(m["content"] for m in messages[:i])
|
||||||
|
ctx_ids = tokenizer.encode(ctx_text)
|
||||||
|
resp_ids = tokenizer.encode(msg["content"], add_special_tokens=False)
|
||||||
|
ctx_ids, resp_ids = _trim(ctx_ids, resp_ids, max_len)
|
||||||
|
if ctx_ids and resp_ids:
|
||||||
|
turns.append((ctx_ids, resp_ids))
|
||||||
|
if not turns:
|
||||||
|
return None
|
||||||
|
raw_scores = _score_batch(
|
||||||
|
turns, model, device, max_len, sentinel_ids=sentinel_ids, per_token=per_token
|
||||||
)
|
)
|
||||||
full_text = tokenizer.apply_chat_template(
|
valid = [s for s in raw_scores if s is not None and s.get("ifd") is not None]
|
||||||
[
|
if not valid:
|
||||||
{"role": "user", "content": instruction},
|
return {"ifd": None, "ifd_turns": raw_scores}
|
||||||
{"role": "assistant", "content": response},
|
avg = sum(s["ifd"] for s in valid) / len(valid)
|
||||||
],
|
|
||||||
tokenize=False,
|
|
||||||
add_generation_prompt=False,
|
|
||||||
)
|
|
||||||
|
|
||||||
full_ids = tokenizer.encode(full_text)
|
|
||||||
prefix_ids = tokenizer.encode(instr_prefix)
|
|
||||||
resp_ids = tokenizer.encode(response)
|
|
||||||
|
|
||||||
if not resp_ids:
|
|
||||||
return {
|
|
||||||
"L_cond": None,
|
|
||||||
"L_uncond": None,
|
|
||||||
"ifd": None,
|
|
||||||
"error": "empty response",
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(full_ids) > max_len:
|
|
||||||
overflow = len(full_ids) - max_len
|
|
||||||
full_ids = full_ids[overflow:]
|
|
||||||
prefix_len = len(prefix_ids) - overflow
|
|
||||||
prefix_len = max(0, prefix_len)
|
|
||||||
else:
|
|
||||||
prefix_len = len(prefix_ids)
|
|
||||||
|
|
||||||
cond_tensor = torch.tensor([full_ids], device=device, dtype=torch.long)
|
|
||||||
|
|
||||||
with torch.inference_mode():
|
|
||||||
logits_qa = model(cond_tensor)["logits"][0]
|
|
||||||
|
|
||||||
resp_start = prefix_len - 1
|
|
||||||
resp_end = len(full_ids) - 1
|
|
||||||
if resp_end <= resp_start:
|
|
||||||
return {
|
|
||||||
"L_cond": None,
|
|
||||||
"L_uncond": None,
|
|
||||||
"ifd": None,
|
|
||||||
"error": "response truncated entirely",
|
|
||||||
}
|
|
||||||
|
|
||||||
resp_logits = logits_qa[resp_start:resp_end]
|
|
||||||
resp_targets = torch.tensor(full_ids[prefix_len:], device=device, dtype=torch.long)
|
|
||||||
L_cond = F.cross_entropy(resp_logits, resp_targets, reduction="mean").item()
|
|
||||||
|
|
||||||
resp_tensor = torch.tensor([resp_ids], device=device, dtype=torch.long)
|
|
||||||
|
|
||||||
with torch.inference_mode():
|
|
||||||
logits_resp = model(resp_tensor)["logits"][0]
|
|
||||||
|
|
||||||
unp_logits = logits_resp[:-1]
|
|
||||||
unp_targets = resp_tensor[0, 1:]
|
|
||||||
L_uncond = F.cross_entropy(unp_logits, unp_targets, reduction="mean").item()
|
|
||||||
|
|
||||||
ifd = L_cond / L_uncond if L_uncond > 0 else None
|
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"L_cond": round(L_cond, 6),
|
"ifd": avg,
|
||||||
"L_uncond": round(L_uncond, 6),
|
"ifd_detail": valid[0] if len(valid) == 1 else None,
|
||||||
"ifd": round(ifd, 6) if ifd is not None else None,
|
"ifd_turns": raw_scores,
|
||||||
"instr_len": prefix_len,
|
|
||||||
"resp_len": len(resp_ids),
|
|
||||||
"error": None,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
def process_file(
|
def process_file(
|
||||||
param_path: str,
|
param_path,
|
||||||
input_file: str,
|
input_file,
|
||||||
output_file: str,
|
output_file,
|
||||||
instr_key: str,
|
instr_key,
|
||||||
resp_key: str,
|
resp_key,
|
||||||
max_len: int = 2048,
|
max_len=2048,
|
||||||
use_chat_template: bool = False,
|
data_format="plain",
|
||||||
|
batch_size=1,
|
||||||
|
device=None,
|
||||||
|
sentinel_text="\n",
|
||||||
|
per_token=False,
|
||||||
):
|
):
|
||||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
if device is None:
|
||||||
dtype = torch.bfloat16 if device == "cuda" else torch.float32
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
|
dtype = torch.bfloat16 if "cuda" in device else torch.float32
|
||||||
|
|
||||||
model = AutoModel.from_pretrained(param_path)
|
model = AutoModel.from_pretrained(param_path)
|
||||||
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
||||||
model.to(device=device, dtype=dtype)
|
model.to(device=device, dtype=dtype)
|
||||||
model.eval()
|
model.eval()
|
||||||
|
|
||||||
if use_chat_template and tokenizer._chat_template is None:
|
sentinel_ids = _resolve_sentinel_ids(tokenizer, sentinel_text)
|
||||||
raise RuntimeError(
|
|
||||||
"--use_chat_template specified but tokenizer has no chat template. "
|
|
||||||
"Add a chat_template to tokenizer_config.json or omit the flag."
|
|
||||||
)
|
|
||||||
|
|
||||||
with open(input_file, "r", encoding="utf-8") as f:
|
with open(input_file, encoding="utf-8") as f:
|
||||||
data = [json.loads(line) for line in f if line.strip()]
|
data = [json.loads(line) for line in f if line.strip()]
|
||||||
|
|
||||||
results = []
|
results = []
|
||||||
ifd_values = []
|
all_ifds = []
|
||||||
|
buffer = []
|
||||||
|
|
||||||
with torch.inference_mode():
|
for item in tqdm.tqdm(data, desc="Computing IFD", unit="sample"):
|
||||||
for item in tqdm.tqdm(data, desc="Computing IFD", unit="sample"):
|
if data_format == "messages":
|
||||||
instruction = item[instr_key]
|
turns = []
|
||||||
response = item[resp_key]
|
for i, msg in enumerate(item.get("messages", [])):
|
||||||
scores = compute_ifd(
|
if msg.get("role") != "assistant":
|
||||||
|
continue
|
||||||
|
ctx_text = "\n\n".join(m["content"] for m in item["messages"][:i])
|
||||||
|
ctx_ids = tokenizer.encode(ctx_text)
|
||||||
|
resp_ids = tokenizer.encode(msg["content"], add_special_tokens=False)
|
||||||
|
ctx_ids, resp_ids = _trim(ctx_ids, resp_ids, max_len)
|
||||||
|
if ctx_ids and resp_ids:
|
||||||
|
turns.append((ctx_ids, resp_ids))
|
||||||
|
if not turns:
|
||||||
|
results.append(
|
||||||
|
{
|
||||||
|
**item,
|
||||||
|
"ifd": None,
|
||||||
|
"skip_reason": "no valid assistant turns",
|
||||||
|
"ifd_turns": [],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
buffer.append((item, turns, "messages"))
|
||||||
|
else:
|
||||||
|
ctx_ids = tokenizer.encode(item[instr_key], add_special_tokens=False)
|
||||||
|
resp_ids = tokenizer.encode(item[resp_key], add_special_tokens=False)
|
||||||
|
ctx_ids, resp_ids = _trim(ctx_ids, resp_ids, max_len)
|
||||||
|
if not ctx_ids or not resp_ids:
|
||||||
|
results.append(
|
||||||
|
{
|
||||||
|
**item,
|
||||||
|
"ifd": None,
|
||||||
|
"ifd_detail": {"skip_reason": "empty ctx or resp"},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
buffer.append((item, [(ctx_ids, resp_ids)], "plain"))
|
||||||
|
|
||||||
|
if len(buffer) >= batch_size:
|
||||||
|
_flush_buffer(
|
||||||
|
buffer,
|
||||||
|
results,
|
||||||
|
all_ifds,
|
||||||
model,
|
model,
|
||||||
tokenizer,
|
|
||||||
instruction,
|
|
||||||
response,
|
|
||||||
device,
|
device,
|
||||||
max_len,
|
max_len,
|
||||||
use_chat_template=use_chat_template,
|
sentinel_ids,
|
||||||
|
per_token,
|
||||||
)
|
)
|
||||||
ifd_values.append(scores["ifd"])
|
|
||||||
results.append({**item, "ifd": scores["ifd"], "ifd_detail": scores})
|
if buffer:
|
||||||
|
_flush_buffer(
|
||||||
|
buffer, results, all_ifds, model, device, max_len, sentinel_ids, per_token
|
||||||
|
)
|
||||||
|
|
||||||
with open(output_file, "w", encoding="utf-8") as f:
|
with open(output_file, "w", encoding="utf-8") as f:
|
||||||
for item in results:
|
for item in results:
|
||||||
f.write(json.dumps(item, ensure_ascii=False) + "\n")
|
f.write(json.dumps(item, ensure_ascii=False) + "\n")
|
||||||
|
|
||||||
valid_ifd = [v for v in ifd_values if v is not None]
|
valid_ifd = [v for v in all_ifds if v is not None]
|
||||||
if valid_ifd:
|
if valid_ifd:
|
||||||
import statistics
|
|
||||||
|
|
||||||
print(f"\n{'=' * 50}")
|
print(f"\n{'=' * 50}")
|
||||||
print(f" Samples: {len(data)}")
|
print(f" Samples: {len(data)}")
|
||||||
print(f" Valid IFD: {len(valid_ifd)}")
|
print(f" Valid IFD: {len(valid_ifd)}")
|
||||||
print(f" Mean IFD: {statistics.mean(valid_ifd):.4f}")
|
print(f" Skipped: {len(data) - len(valid_ifd)}")
|
||||||
print(f" Median IFD: {statistics.median(valid_ifd):.4f}")
|
print(f" Mean IFD: {statistics.mean(valid_ifd):.4f}")
|
||||||
print(f" Stdev IFD: {statistics.stdev(valid_ifd):.4f}")
|
print(f" Median IFD: {statistics.median(valid_ifd):.4f}")
|
||||||
print(f" Min IFD: {min(valid_ifd):.4f}")
|
if len(valid_ifd) > 1:
|
||||||
print(f" Max IFD: {max(valid_ifd):.4f}")
|
print(f" Stdev IFD: {statistics.stdev(valid_ifd):.4f}")
|
||||||
|
print(f" Min IFD: {min(valid_ifd):.4f}")
|
||||||
|
print(f" Max IFD: {max(valid_ifd):.4f}")
|
||||||
print(f"{'=' * 50}")
|
print(f"{'=' * 50}")
|
||||||
|
|
||||||
print(f"Results saved to {output_file}")
|
print(f"Results saved to {output_file}")
|
||||||
|
|
||||||
|
|
||||||
|
def _flush_buffer(
|
||||||
|
buffer, results, all_ifds, model, device, max_len, sentinel_ids, per_token
|
||||||
|
):
|
||||||
|
all_pairs = []
|
||||||
|
indices = []
|
||||||
|
for item, turns, fmt in buffer:
|
||||||
|
start = len(all_pairs)
|
||||||
|
all_pairs.extend(turns)
|
||||||
|
indices.append((item, turns, fmt, start, len(all_pairs)))
|
||||||
|
|
||||||
|
raw = _score_batch(
|
||||||
|
all_pairs,
|
||||||
|
model,
|
||||||
|
device,
|
||||||
|
max_len,
|
||||||
|
sentinel_ids=sentinel_ids,
|
||||||
|
per_token=per_token,
|
||||||
|
)
|
||||||
|
|
||||||
|
for item, turns, fmt, start, end in indices:
|
||||||
|
turn_scores = raw[start:end]
|
||||||
|
if fmt == "messages":
|
||||||
|
valid = [
|
||||||
|
s for s in turn_scores if s is not None and s.get("ifd") is not None
|
||||||
|
]
|
||||||
|
if not valid:
|
||||||
|
results.append({**item, "ifd": None, "ifd_turns": turn_scores})
|
||||||
|
else:
|
||||||
|
avg = sum(s["ifd"] for s in valid) / len(valid)
|
||||||
|
all_ifds.append(avg)
|
||||||
|
results.append(
|
||||||
|
{
|
||||||
|
**item,
|
||||||
|
"ifd": avg,
|
||||||
|
"ifd_detail": valid[0] if len(valid) == 1 else None,
|
||||||
|
"ifd_turns": turn_scores,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
score = turn_scores[0]
|
||||||
|
all_ifds.append(score.get("ifd"))
|
||||||
|
results.append({**item, "ifd": score.get("ifd"), "ifd_detail": score})
|
||||||
|
|
||||||
|
buffer.clear()
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
parser = argparse.ArgumentParser(
|
parser = argparse.ArgumentParser(
|
||||||
description="Compute IFD scores for instruction-response data"
|
description="Compute IFD scores for instruction-response data"
|
||||||
@@ -266,29 +424,34 @@ def main():
|
|||||||
parser.add_argument("--param_path", type=str, required=True, help="Model directory")
|
parser.add_argument("--param_path", type=str, required=True, help="Model directory")
|
||||||
parser.add_argument("--input", type=str, required=True, help="Input JSONL file")
|
parser.add_argument("--input", type=str, required=True, help="Input JSONL file")
|
||||||
parser.add_argument("--output", type=str, required=True, help="Output JSONL file")
|
parser.add_argument("--output", type=str, required=True, help="Output JSONL file")
|
||||||
|
parser.add_argument("--max_len", type=int, default=2048, help="Max token length")
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--instr_key",
|
"--format",
|
||||||
type=str,
|
type=str,
|
||||||
default="instruction",
|
default="plain",
|
||||||
help="Key for instruction field",
|
choices=["plain", "messages"],
|
||||||
|
help="Input format",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--resp_key",
|
"--instr_key", type=str, default="instruction", help="Key for instruction field"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--resp_key", type=str, default="response", help="Key for response field"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--batch_size", type=int, default=8, help="Batch size for model forward passes"
|
||||||
|
)
|
||||||
|
parser.add_argument("--device", type=str, default=None, help="Device (e.g. cuda:0)")
|
||||||
|
parser.add_argument(
|
||||||
|
"--sentinel_text",
|
||||||
type=str,
|
type=str,
|
||||||
default="response",
|
default="\n",
|
||||||
help="Key for response field",
|
help='Plain-text prefix for unconditional pass (default: "\\n"). Use "" for bos/pad fallback.',
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--max_len",
|
"--per_token",
|
||||||
type=int,
|
|
||||||
default=2048,
|
|
||||||
help="Max token length (instruction truncated to fit)",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--no_chat_template",
|
|
||||||
action="store_true",
|
action="store_true",
|
||||||
default=False,
|
help="Include per-token IFD breakdown in output",
|
||||||
help="Disable chat template, use raw text concatenation",
|
|
||||||
)
|
)
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
@@ -299,7 +462,11 @@ def main():
|
|||||||
args.instr_key,
|
args.instr_key,
|
||||||
args.resp_key,
|
args.resp_key,
|
||||||
args.max_len,
|
args.max_len,
|
||||||
use_chat_template=not args.no_chat_template,
|
data_format=args.format,
|
||||||
|
batch_size=args.batch_size,
|
||||||
|
device=args.device,
|
||||||
|
sentinel_text=args.sentinel_text,
|
||||||
|
per_token=args.per_token,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,153 @@
|
|||||||
|
"""ROUGE evaluation (manual implementation, no external deps).
|
||||||
|
|
||||||
|
Computes ROUGE-1, ROUGE-2, ROUGE-L precision, recall, and F1.
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
# Batch evaluation from JSONL (each line: {"reference": ..., "candidate": ...})
|
||||||
|
python scripts/eval/evaluate_rouge.py --data_path preds.jsonl --output results.json
|
||||||
|
|
||||||
|
# As a library
|
||||||
|
from scripts.eval.evaluate_rouge import compute_rouge
|
||||||
|
scores = compute_rouge("the cat sat on the mat", "the cat sat")
|
||||||
|
"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import json
|
||||||
|
from collections import Counter
|
||||||
|
from typing import Dict, List, Tuple
|
||||||
|
|
||||||
|
|
||||||
|
def _tokenize(text: str) -> List[str]:
|
||||||
|
return text.split()
|
||||||
|
|
||||||
|
|
||||||
|
def _ngrams(tokens: List[str], n: int) -> Counter:
|
||||||
|
return Counter(zip(*[tokens[i:] for i in range(n)]))
|
||||||
|
|
||||||
|
|
||||||
|
def _lcs(x: List[str], y: List[str]) -> int:
|
||||||
|
m, n = len(x), len(y)
|
||||||
|
dp = [[0] * (n + 1) for _ in range(m + 1)]
|
||||||
|
for i in range(1, m + 1):
|
||||||
|
xi = x[i - 1]
|
||||||
|
dpi = dp[i]
|
||||||
|
dpi_1 = dp[i - 1]
|
||||||
|
for j in range(1, n + 1):
|
||||||
|
if xi == y[j - 1]:
|
||||||
|
dpi[j] = dpi_1[j - 1] + 1
|
||||||
|
else:
|
||||||
|
dpi[j] = dpi_1[j] if dpi_1[j] > dpi[j - 1] else dpi[j - 1]
|
||||||
|
return dp[m][n]
|
||||||
|
|
||||||
|
|
||||||
|
def _f1(precision: float, recall: float) -> float:
|
||||||
|
if precision + recall == 0:
|
||||||
|
return 0.0
|
||||||
|
return 2 * precision * recall / (precision + recall)
|
||||||
|
|
||||||
|
|
||||||
|
def _rouge_n(ref_tokens: List[str], cand_tokens: List[str], n: int) -> Dict[str, float]:
|
||||||
|
ref_ngrams = _ngrams(ref_tokens, n)
|
||||||
|
cand_ngrams = _ngrams(cand_tokens, n)
|
||||||
|
|
||||||
|
overlap = sum((cand_ngrams & ref_ngrams).values())
|
||||||
|
cand_total = sum(cand_ngrams.values())
|
||||||
|
ref_total = sum(ref_ngrams.values())
|
||||||
|
|
||||||
|
precision = overlap / cand_total if cand_total > 0 else 0.0
|
||||||
|
recall = overlap / ref_total if ref_total > 0 else 0.0
|
||||||
|
f1 = _f1(precision, recall)
|
||||||
|
|
||||||
|
return {"precision": precision, "recall": recall, "f1": f1}
|
||||||
|
|
||||||
|
|
||||||
|
def _rouge_l(ref_tokens: List[str], cand_tokens: List[str]) -> Dict[str, float]:
|
||||||
|
lcs_len = _lcs(ref_tokens, cand_tokens)
|
||||||
|
ref_len = len(ref_tokens)
|
||||||
|
cand_len = len(cand_tokens)
|
||||||
|
|
||||||
|
recall = lcs_len / ref_len if ref_len > 0 else 0.0
|
||||||
|
precision = lcs_len / cand_len if cand_len > 0 else 0.0
|
||||||
|
f1 = _f1(precision, recall)
|
||||||
|
|
||||||
|
return {"precision": precision, "recall": recall, "f1": f1}
|
||||||
|
|
||||||
|
|
||||||
|
def compute_rouge(
|
||||||
|
reference: str, candidate: str, n: int = 2
|
||||||
|
) -> Dict[str, Dict[str, float]]:
|
||||||
|
"""Compute ROUGE-N (1..n) and ROUGE-L scores.
|
||||||
|
|
||||||
|
Returns::
|
||||||
|
|
||||||
|
{
|
||||||
|
"rouge-1": {"precision": ..., "recall": ..., "f1": ...},
|
||||||
|
"rouge-2": {"precision": ..., "recall": ..., "f1": ...},
|
||||||
|
"rouge-l": {"precision": ..., "recall": ..., "f1": ...},
|
||||||
|
}
|
||||||
|
"""
|
||||||
|
ref_tokens = _tokenize(reference)
|
||||||
|
cand_tokens = _tokenize(candidate)
|
||||||
|
|
||||||
|
results = {}
|
||||||
|
for i in range(1, n + 1):
|
||||||
|
results[f"rouge-{i}"] = _rouge_n(ref_tokens, cand_tokens, i)
|
||||||
|
results["rouge-l"] = _rouge_l(ref_tokens, cand_tokens)
|
||||||
|
return results
|
||||||
|
|
||||||
|
|
||||||
|
def evaluate_file(data_path: str) -> Dict:
|
||||||
|
with open(data_path, "r", encoding="utf-8") as f:
|
||||||
|
pairs = [json.loads(line) for line in f if line.strip()]
|
||||||
|
|
||||||
|
agg = {
|
||||||
|
k: {"precision": 0.0, "recall": 0.0, "f1": 0.0}
|
||||||
|
for k in ("rouge-1", "rouge-2", "rouge-l")
|
||||||
|
}
|
||||||
|
per_item = []
|
||||||
|
|
||||||
|
for item in pairs:
|
||||||
|
ref = item["reference"]
|
||||||
|
cand = item["candidate"]
|
||||||
|
scores = compute_rouge(ref, cand)
|
||||||
|
per_item.append({**item, "scores": scores})
|
||||||
|
for k, v in scores.items():
|
||||||
|
agg[k]["precision"] += v["precision"]
|
||||||
|
agg[k]["recall"] += v["recall"]
|
||||||
|
agg[k]["f1"] += v["f1"]
|
||||||
|
|
||||||
|
n = len(pairs)
|
||||||
|
for k in agg:
|
||||||
|
agg[k] = {m: v / n for m, v in agg[k].items()}
|
||||||
|
|
||||||
|
return {"num_samples": n, "aggregate": agg, "per_item": per_item}
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
parser = argparse.ArgumentParser(description="ROUGE evaluation")
|
||||||
|
parser.add_argument(
|
||||||
|
"--data_path", required=True, help="JSONL with reference/candidate per line"
|
||||||
|
)
|
||||||
|
parser.add_argument("--output", type=str, default=None, help="Output JSON path")
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
results = evaluate_file(args.data_path)
|
||||||
|
agg = results["aggregate"]
|
||||||
|
|
||||||
|
print(f"Samples: {results['num_samples']}")
|
||||||
|
print()
|
||||||
|
for metric in ("rouge-1", "rouge-2", "rouge-l"):
|
||||||
|
s = agg[metric]
|
||||||
|
print(
|
||||||
|
f" {metric:8s} P={s['precision']:.4f} R={s['recall']:.4f} F1={s['f1']:.4f}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if args.output:
|
||||||
|
with open(args.output, "w", encoding="utf-8") as f:
|
||||||
|
json.dump(results, f, indent=2, ensure_ascii=False)
|
||||||
|
print(f"\nSaved to {args.output}")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
+129
-79
@@ -1,12 +1,13 @@
|
|||||||
"""Benchmark AutoRegressiveLM with KVCache"""
|
"""Benchmark AutoRegressiveLM with KVCache"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, Dict
|
from typing import Any, Dict
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from astrai.config import AutoRegressiveLMConfig
|
from astrai.config import AutoRegressiveLMConfig
|
||||||
from astrai.inference import KVCache
|
from astrai.inference import ContiguousCache, PageCache
|
||||||
from astrai.model.transformer import AutoRegressiveLM
|
from astrai.model.transformer import AutoRegressiveLM
|
||||||
|
|
||||||
|
|
||||||
@@ -24,41 +25,14 @@ class GenerationBenchmark:
|
|||||||
config: AutoRegressiveLMConfig,
|
config: AutoRegressiveLMConfig,
|
||||||
device: str = "cuda",
|
device: str = "cuda",
|
||||||
dtype: torch.dtype = torch.bfloat16,
|
dtype: torch.dtype = torch.bfloat16,
|
||||||
page_size: int = 128,
|
cache_type: str = "contiguous",
|
||||||
):
|
):
|
||||||
self.config = config
|
self.config = config
|
||||||
self.device = device
|
self.device = device
|
||||||
self.dtype = dtype
|
self.dtype = dtype
|
||||||
|
self.cache_type = cache_type
|
||||||
self.model = AutoRegressiveLM(config).to(device=device, dtype=dtype)
|
self.model = AutoRegressiveLM(config).to(device=device, dtype=dtype)
|
||||||
self.model.eval()
|
self.model.eval()
|
||||||
head_dim = config.dim // config.n_heads
|
|
||||||
n_pages = (config.max_len * 4 + page_size - 1) // page_size
|
|
||||||
self._page_cache = KVCache(
|
|
||||||
config.n_layers,
|
|
||||||
n_pages,
|
|
||||||
page_size,
|
|
||||||
config.n_kv_heads,
|
|
||||||
head_dim,
|
|
||||||
device,
|
|
||||||
dtype,
|
|
||||||
)
|
|
||||||
|
|
||||||
def _prepare_inputs(self, batch_size: int, prompt_length: int, total_length: int):
|
|
||||||
prompt_ids = torch.randint(
|
|
||||||
low=0,
|
|
||||||
high=self.config.vocab_size,
|
|
||||||
size=(batch_size, prompt_length),
|
|
||||||
device=self.device,
|
|
||||||
dtype=torch.long,
|
|
||||||
)
|
|
||||||
gen_ids = torch.randint(
|
|
||||||
low=0,
|
|
||||||
high=self.config.vocab_size,
|
|
||||||
size=(batch_size, total_length - prompt_length),
|
|
||||||
device=self.device,
|
|
||||||
dtype=torch.long,
|
|
||||||
)
|
|
||||||
return prompt_ids, gen_ids
|
|
||||||
|
|
||||||
@torch.inference_mode()
|
@torch.inference_mode()
|
||||||
def run_prefill_benchmark(
|
def run_prefill_benchmark(
|
||||||
@@ -68,8 +42,12 @@ class GenerationBenchmark:
|
|||||||
num_trials: int = 10,
|
num_trials: int = 10,
|
||||||
) -> BenchmarkResult:
|
) -> BenchmarkResult:
|
||||||
for _ in range(3):
|
for _ in range(3):
|
||||||
prompt_ids, _ = self._prepare_inputs(
|
prompt_ids = torch.randint(
|
||||||
batch_size, prompt_length, prompt_length
|
0,
|
||||||
|
self.config.vocab_size,
|
||||||
|
(batch_size, prompt_length),
|
||||||
|
device=self.device,
|
||||||
|
dtype=torch.long,
|
||||||
)
|
)
|
||||||
_ = self.model(prompt_ids)
|
_ = self.model(prompt_ids)
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
@@ -78,12 +56,15 @@ class GenerationBenchmark:
|
|||||||
total_tokens = batch_size * prompt_length * num_trials
|
total_tokens = batch_size * prompt_length * num_trials
|
||||||
|
|
||||||
for trial in range(num_trials):
|
for trial in range(num_trials):
|
||||||
prompt_ids, _ = self._prepare_inputs(
|
prompt_ids = torch.randint(
|
||||||
batch_size, prompt_length, prompt_length
|
0,
|
||||||
|
self.config.vocab_size,
|
||||||
|
(batch_size, prompt_length),
|
||||||
|
device=self.device,
|
||||||
|
dtype=torch.long,
|
||||||
)
|
)
|
||||||
start = torch.cuda.Event(enable_timing=True)
|
start = torch.cuda.Event(enable_timing=True)
|
||||||
end = torch.cuda.Event(enable_timing=True)
|
end = torch.cuda.Event(enable_timing=True)
|
||||||
|
|
||||||
start.record()
|
start.record()
|
||||||
_ = self.model(prompt_ids)
|
_ = self.model(prompt_ids)
|
||||||
end.record()
|
end.record()
|
||||||
@@ -107,6 +88,7 @@ class GenerationBenchmark:
|
|||||||
"prompt_length": prompt_length,
|
"prompt_length": prompt_length,
|
||||||
"dtype": str(self.dtype),
|
"dtype": str(self.dtype),
|
||||||
"device": self.device,
|
"device": self.device,
|
||||||
|
"cache": "none",
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -120,29 +102,56 @@ class GenerationBenchmark:
|
|||||||
) -> BenchmarkResult:
|
) -> BenchmarkResult:
|
||||||
total_time = 0.0
|
total_time = 0.0
|
||||||
total_tokens = batch_size * gen_length * num_trials
|
total_tokens = batch_size * gen_length * num_trials
|
||||||
page_size = self._page_cache.page_size
|
|
||||||
|
|
||||||
for trial in range(num_trials):
|
for trial in range(num_trials):
|
||||||
prompt_ids, gen_ids = self._prepare_inputs(
|
prompt_ids = torch.randint(
|
||||||
batch_size,
|
0,
|
||||||
prompt_length,
|
self.config.vocab_size,
|
||||||
prompt_length + gen_length,
|
(batch_size, prompt_length),
|
||||||
)
|
|
||||||
|
|
||||||
n_pages = (prompt_length + gen_length + page_size - 1) // page_size
|
|
||||||
total = n_pages * batch_size
|
|
||||||
pages = []
|
|
||||||
for _ in range(total):
|
|
||||||
p = self._page_cache._pool.alloc()
|
|
||||||
assert p >= 0, "OOM"
|
|
||||||
pages.append(p)
|
|
||||||
page_table = torch.tensor(
|
|
||||||
[pages[i * n_pages : (i + 1) * n_pages] for i in range(batch_size)],
|
|
||||||
dtype=torch.long,
|
|
||||||
device=self.device,
|
device=self.device,
|
||||||
|
dtype=torch.long,
|
||||||
|
)
|
||||||
|
gen_ids = torch.randint(
|
||||||
|
0,
|
||||||
|
self.config.vocab_size,
|
||||||
|
(batch_size, gen_length),
|
||||||
|
device=self.device,
|
||||||
|
dtype=torch.long,
|
||||||
)
|
)
|
||||||
|
|
||||||
cv = self._page_cache.bind(page_table, total_len=prompt_length)
|
head_dim = self.config.dim // self.config.n_heads
|
||||||
|
max_seq = prompt_length + gen_length
|
||||||
|
|
||||||
|
if self.cache_type == "contiguous":
|
||||||
|
cache = ContiguousCache(
|
||||||
|
self.config.n_layers,
|
||||||
|
batch_size,
|
||||||
|
max_seq,
|
||||||
|
self.config.n_kv_heads,
|
||||||
|
head_dim,
|
||||||
|
self.device,
|
||||||
|
self.dtype,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
page_size = 128
|
||||||
|
n_pages = (max_seq + page_size - 1) // page_size * batch_size
|
||||||
|
cache = PageCache(
|
||||||
|
self.config.n_layers,
|
||||||
|
n_pages,
|
||||||
|
page_size,
|
||||||
|
self.config.n_kv_heads,
|
||||||
|
head_dim,
|
||||||
|
self.device,
|
||||||
|
self.dtype,
|
||||||
|
)
|
||||||
|
|
||||||
|
task_ids = [f"b{i}" for i in range(batch_size)]
|
||||||
|
for tid in task_ids:
|
||||||
|
cache.task_alloc(tid, [0] * max_seq)
|
||||||
|
for p in range(max_seq):
|
||||||
|
cache.task_extend(tid, p)
|
||||||
|
|
||||||
|
cv = cache.bind_tasks(task_ids, prompt_length, self.device)
|
||||||
_ = self.model(
|
_ = self.model(
|
||||||
prompt_ids,
|
prompt_ids,
|
||||||
paged_cache=cv,
|
paged_cache=cv,
|
||||||
@@ -152,37 +161,35 @@ class GenerationBenchmark:
|
|||||||
.unsqueeze(0)
|
.unsqueeze(0)
|
||||||
.expand(batch_size, -1),
|
.expand(batch_size, -1),
|
||||||
)
|
)
|
||||||
|
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
|
|
||||||
start = torch.cuda.Event(enable_timing=True)
|
start = torch.cuda.Event(enable_timing=True)
|
||||||
end = torch.cuda.Event(enable_timing=True)
|
end = torch.cuda.Event(enable_timing=True)
|
||||||
|
|
||||||
start.record()
|
start.record()
|
||||||
current_pos = prompt_length
|
|
||||||
for i in range(gen_length):
|
for i in range(gen_length):
|
||||||
input_token = gen_ids[:, i : i + 1]
|
pos = prompt_length + i
|
||||||
cv = self._page_cache.bind(page_table, total_len=current_pos + 1)
|
cv = cache.bind_tasks(task_ids, pos + 1, self.device)
|
||||||
_ = self.model(
|
_ = self.model(
|
||||||
input_token,
|
gen_ids[:, i : i + 1],
|
||||||
paged_cache=cv,
|
paged_cache=cv,
|
||||||
position_ids=torch.full(
|
position_ids=torch.full(
|
||||||
(batch_size, 1),
|
(batch_size, 1),
|
||||||
current_pos,
|
pos,
|
||||||
dtype=torch.long,
|
dtype=torch.long,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
current_pos += 1
|
|
||||||
end.record()
|
end.record()
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
|
|
||||||
|
for tid in task_ids:
|
||||||
|
cache.task_free(tid)
|
||||||
|
|
||||||
trial_time = start.elapsed_time(end) / 1000
|
trial_time = start.elapsed_time(end) / 1000
|
||||||
total_time += trial_time
|
total_time += trial_time
|
||||||
|
|
||||||
for idx in pages:
|
|
||||||
self._page_cache._pool.free(idx)
|
|
||||||
|
|
||||||
print(
|
print(
|
||||||
f" Trial {trial + 1}/{num_trials}: {gen_length} tokens in {trial_time:.3f}s "
|
f" Trial {trial + 1}/{num_trials}: {gen_length} tokens in {trial_time:.3f}s "
|
||||||
f"({gen_length / trial_time:.1f} tok/s)"
|
f"({gen_length / trial_time:.1f} tok/s)"
|
||||||
@@ -199,6 +206,7 @@ class GenerationBenchmark:
|
|||||||
"gen_length": gen_length,
|
"gen_length": gen_length,
|
||||||
"dtype": str(self.dtype),
|
"dtype": str(self.dtype),
|
||||||
"device": self.device,
|
"device": self.device,
|
||||||
|
"cache": self.cache_type,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -216,6 +224,42 @@ def print_benchmark_result(result: BenchmarkResult):
|
|||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
parser = argparse.ArgumentParser(description="AutoRegressiveLM benchmark")
|
||||||
|
parser.add_argument(
|
||||||
|
"--device", type=str, default="cuda", help="Device (default: cuda)"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--dtype",
|
||||||
|
type=str,
|
||||||
|
default="bfloat16",
|
||||||
|
choices=["bfloat16", "float16", "float32"],
|
||||||
|
help="Dtype",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--cache",
|
||||||
|
type=str,
|
||||||
|
default="contiguous",
|
||||||
|
choices=["contiguous", "paged"],
|
||||||
|
help="KV cache type",
|
||||||
|
)
|
||||||
|
parser.add_argument("--batch_size", type=int, default=4, help="Batch size")
|
||||||
|
parser.add_argument("--prompt_length", type=int, default=512, help="Prompt length")
|
||||||
|
parser.add_argument("--gen_length", type=int, default=128, help="Generation length")
|
||||||
|
parser.add_argument("--num_trials", type=int, default=5, help="Number of trials")
|
||||||
|
parser.add_argument(
|
||||||
|
"--prefill_only", action="store_true", help="Run prefill benchmark only"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--decode_only", action="store_true", help="Run decoding benchmark only"
|
||||||
|
)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
dtype_map = {
|
||||||
|
"bfloat16": torch.bfloat16,
|
||||||
|
"float16": torch.float16,
|
||||||
|
"float32": torch.float32,
|
||||||
|
}
|
||||||
|
|
||||||
config = AutoRegressiveLMConfig(
|
config = AutoRegressiveLMConfig(
|
||||||
vocab_size=10000,
|
vocab_size=10000,
|
||||||
dim=1536,
|
dim=1536,
|
||||||
@@ -227,23 +271,29 @@ if __name__ == "__main__":
|
|||||||
norm_eps=1e-5,
|
norm_eps=1e-5,
|
||||||
)
|
)
|
||||||
|
|
||||||
benchmark = GenerationBenchmark(config)
|
benchmark = GenerationBenchmark(
|
||||||
|
config, device=args.device, dtype=dtype_map[args.dtype], cache_type=args.cache
|
||||||
|
)
|
||||||
|
|
||||||
print("=" * 80)
|
print("=" * 80)
|
||||||
print("Running AutoRegressiveLM Generation Benchmark (KVCache)")
|
print(
|
||||||
|
f"Running AutoRegressiveLM Benchmark (device={args.device}, dtype={args.dtype})"
|
||||||
|
)
|
||||||
print("=" * 80)
|
print("=" * 80)
|
||||||
|
|
||||||
prefill_result = benchmark.run_prefill_benchmark(
|
if not args.decode_only:
|
||||||
batch_size=4,
|
prefill_result = benchmark.run_prefill_benchmark(
|
||||||
prompt_length=512,
|
batch_size=args.batch_size,
|
||||||
num_trials=5,
|
prompt_length=args.prompt_length,
|
||||||
)
|
num_trials=args.num_trials,
|
||||||
print_benchmark_result(prefill_result)
|
)
|
||||||
|
print_benchmark_result(prefill_result)
|
||||||
|
|
||||||
gen_result = benchmark.run_decoding_benchmark(
|
if not args.prefill_only:
|
||||||
batch_size=4,
|
gen_result = benchmark.run_decoding_benchmark(
|
||||||
prompt_length=512,
|
batch_size=args.batch_size,
|
||||||
gen_length=128,
|
prompt_length=args.prompt_length,
|
||||||
num_trials=5,
|
gen_length=args.gen_length,
|
||||||
)
|
num_trials=args.num_trials,
|
||||||
print_benchmark_result(gen_result)
|
)
|
||||||
|
print_benchmark_result(gen_result)
|
||||||
|
|||||||
+117
-45
@@ -1,9 +1,11 @@
|
|||||||
import argparse
|
import argparse
|
||||||
import os
|
import os
|
||||||
from functools import partial
|
from functools import partial
|
||||||
|
from typing import Any, Dict
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.optim as optim
|
import torch.optim as optim
|
||||||
|
from torch import Tensor, nn
|
||||||
|
|
||||||
from astrai.config import AutoRegressiveLMConfig, TrainConfig
|
from astrai.config import AutoRegressiveLMConfig, TrainConfig
|
||||||
from astrai.dataset import DatasetFactory
|
from astrai.dataset import DatasetFactory
|
||||||
@@ -12,6 +14,84 @@ from astrai.model.components.decoder_block import DecoderBlock
|
|||||||
from astrai.trainer import SchedulerFactory, Trainer
|
from astrai.trainer import SchedulerFactory, Trainer
|
||||||
|
|
||||||
|
|
||||||
|
class MuonMix(optim.Optimizer):
|
||||||
|
"""Combined Muon (matrix) + AdamW (non-matrix) optimizer."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model: nn.Module,
|
||||||
|
lr: float = 3e-4,
|
||||||
|
weight_decay: float = 0.1,
|
||||||
|
momentum: float = 0.95,
|
||||||
|
nesterov: bool = True,
|
||||||
|
ns_steps: int = 5,
|
||||||
|
adjust_lr_fn: str = "match_rms_adamw",
|
||||||
|
):
|
||||||
|
defaults = dict(
|
||||||
|
lr=lr,
|
||||||
|
weight_decay=weight_decay,
|
||||||
|
momentum=momentum,
|
||||||
|
nesterov=nesterov,
|
||||||
|
ns_steps=ns_steps,
|
||||||
|
adjust_lr_fn=adjust_lr_fn,
|
||||||
|
)
|
||||||
|
params = [p for p in model.parameters() if p.requires_grad]
|
||||||
|
super().__init__(params, defaults)
|
||||||
|
|
||||||
|
matrix_params: list[Tensor] = []
|
||||||
|
other_params: list[Tensor] = []
|
||||||
|
for name, param in model.named_parameters():
|
||||||
|
if not param.requires_grad:
|
||||||
|
continue
|
||||||
|
if (
|
||||||
|
param.dim() >= 2
|
||||||
|
and "norm" not in name
|
||||||
|
and "bias" not in name
|
||||||
|
and "embed" not in name
|
||||||
|
and "lm_head" not in name
|
||||||
|
):
|
||||||
|
matrix_params.append(param)
|
||||||
|
else:
|
||||||
|
other_params.append(param)
|
||||||
|
|
||||||
|
self.muon = optim.Muon(
|
||||||
|
matrix_params,
|
||||||
|
lr=lr,
|
||||||
|
weight_decay=weight_decay,
|
||||||
|
momentum=momentum,
|
||||||
|
nesterov=nesterov,
|
||||||
|
ns_steps=ns_steps,
|
||||||
|
adjust_lr_fn=adjust_lr_fn,
|
||||||
|
)
|
||||||
|
self.adamw = optim.AdamW(
|
||||||
|
[{"params": other_params, "weight_decay": 0.0}],
|
||||||
|
lr=lr,
|
||||||
|
betas=(0.9, 0.95),
|
||||||
|
fused=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.param_groups = [*self.muon.param_groups, *self.adamw.param_groups]
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def step(self, closure=None):
|
||||||
|
self.muon.step(closure)
|
||||||
|
self.adamw.step(closure)
|
||||||
|
|
||||||
|
def zero_grad(self, set_to_none: bool = True):
|
||||||
|
self.muon.zero_grad(set_to_none)
|
||||||
|
self.adamw.zero_grad(set_to_none)
|
||||||
|
|
||||||
|
def state_dict(self) -> Dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"muon": self.muon.state_dict(),
|
||||||
|
"adamw": self.adamw.state_dict(),
|
||||||
|
}
|
||||||
|
|
||||||
|
def load_state_dict(self, state_dict: Dict[str, Any]):
|
||||||
|
self.muon.load_state_dict(state_dict["muon"])
|
||||||
|
self.adamw.load_state_dict(state_dict["adamw"])
|
||||||
|
|
||||||
|
|
||||||
def parse_args() -> argparse.Namespace:
|
def parse_args() -> argparse.Namespace:
|
||||||
|
|
||||||
parser = argparse.ArgumentParser(description="Train the AutoRegressiveLM model.")
|
parser = argparse.ArgumentParser(description="Train the AutoRegressiveLM model.")
|
||||||
@@ -64,22 +144,35 @@ def parse_args() -> argparse.Namespace:
|
|||||||
help="Max gradient norm for clipping.",
|
help="Max gradient norm for clipping.",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--adamw_beta1",
|
"--weight_decay",
|
||||||
type=float,
|
type=float,
|
||||||
default=0.9,
|
default=0.1,
|
||||||
help="Beta1 for AdamW optimizer.",
|
help="Weight decay (applied to Muon matrix params; non-matrix use 0).",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--adamw_beta2",
|
"--muon_momentum",
|
||||||
type=float,
|
type=float,
|
||||||
default=0.95,
|
default=0.95,
|
||||||
help="Beta2 for AdamW optimizer.",
|
help="Momentum factor for Muon optimizer.",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--adamw_weight_decay",
|
"--muon_nesterov",
|
||||||
type=float,
|
action=argparse.BooleanOptionalAction,
|
||||||
default=0.01,
|
default=True,
|
||||||
help="Weight decay for AdamW optimizer.",
|
help="Enable Nesterov momentum for Muon.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--muon_ns_steps",
|
||||||
|
type=int,
|
||||||
|
default=5,
|
||||||
|
help="Newton-Schulz iteration steps for Muon.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--muon_adjust_lr",
|
||||||
|
type=str,
|
||||||
|
default="match_rms_adamw",
|
||||||
|
choices=["original", "match_rms_adamw"],
|
||||||
|
help="Muon learning rate adjustment strategy.",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--random_seed", type=int, default=3407, help="Random seed for reproducibility."
|
"--random_seed", type=int, default=3407, help="Random seed for reproducibility."
|
||||||
@@ -265,21 +358,8 @@ def create_model(config):
|
|||||||
return AutoRegressiveLM(config).to(dtype=torch.bfloat16)
|
return AutoRegressiveLM(config).to(dtype=torch.bfloat16)
|
||||||
|
|
||||||
|
|
||||||
def create_optimizer(model, **kwargs) -> optim.Optimizer:
|
def create_optimizer(model, **kwargs) -> MuonMix:
|
||||||
decay_params = []
|
return MuonMix(model, **kwargs)
|
||||||
no_decay_params = []
|
|
||||||
for name, param in model.named_parameters():
|
|
||||||
if not param.requires_grad:
|
|
||||||
continue
|
|
||||||
if param.dim() < 2 or "norm" in name or "bias" in name:
|
|
||||||
no_decay_params.append(param)
|
|
||||||
else:
|
|
||||||
decay_params.append(param)
|
|
||||||
param_groups = [
|
|
||||||
{"params": decay_params, "weight_decay": kwargs.pop("weight_decay", 0.01)},
|
|
||||||
{"params": no_decay_params, "weight_decay": 0.0},
|
|
||||||
]
|
|
||||||
return optim.AdamW(param_groups, fused=True, **kwargs)
|
|
||||||
|
|
||||||
|
|
||||||
def create_scheduler(
|
def create_scheduler(
|
||||||
@@ -310,7 +390,6 @@ def train(
|
|||||||
train_type: str,
|
train_type: str,
|
||||||
param_path: str,
|
param_path: str,
|
||||||
data_root_path: str,
|
data_root_path: str,
|
||||||
max_lr: float,
|
|
||||||
n_epoch: int,
|
n_epoch: int,
|
||||||
batch_per_device: int,
|
batch_per_device: int,
|
||||||
start_epoch: int,
|
start_epoch: int,
|
||||||
@@ -323,16 +402,7 @@ def train(
|
|||||||
val_step: int,
|
val_step: int,
|
||||||
metrics: list[str],
|
metrics: list[str],
|
||||||
log_dir: str,
|
log_dir: str,
|
||||||
dpo_beta: float,
|
|
||||||
grpo_clip_eps: float,
|
|
||||||
grpo_kl_coef: float,
|
|
||||||
group_size: int,
|
|
||||||
grpo_sync_interval: int,
|
|
||||||
adamw_beta1: float,
|
|
||||||
adamw_beta2: float,
|
|
||||||
adamw_weight_decay: float,
|
|
||||||
max_grad_norm: float,
|
max_grad_norm: float,
|
||||||
label_smoothing: float,
|
|
||||||
random_seed: int,
|
random_seed: int,
|
||||||
num_workers: int,
|
num_workers: int,
|
||||||
pin_memory: bool,
|
pin_memory: bool,
|
||||||
@@ -353,6 +423,7 @@ def train(
|
|||||||
t_mult: int,
|
t_mult: int,
|
||||||
stable_steps: int,
|
stable_steps: int,
|
||||||
decay_steps: int,
|
decay_steps: int,
|
||||||
|
**kwargs,
|
||||||
):
|
):
|
||||||
assert train_type in ["seq", "sft", "dpo", "grpo"]
|
assert train_type in ["seq", "sft", "dpo", "grpo"]
|
||||||
assert os.path.exists(param_path)
|
assert os.path.exists(param_path)
|
||||||
@@ -368,12 +439,12 @@ def train(
|
|||||||
window_size = config.max_len
|
window_size = config.max_len
|
||||||
|
|
||||||
strategy_kwargs = {
|
strategy_kwargs = {
|
||||||
"beta": dpo_beta,
|
"beta": kwargs.pop("dpo_beta"),
|
||||||
"label_smoothing": label_smoothing,
|
"label_smoothing": kwargs.pop("label_smoothing"),
|
||||||
"clip_eps": grpo_clip_eps,
|
"clip_eps": kwargs.pop("grpo_clip_eps"),
|
||||||
"kl_coef": grpo_kl_coef,
|
"kl_coef": kwargs.pop("grpo_kl_coef"),
|
||||||
"group_size": group_size,
|
"group_size": kwargs.pop("group_size"),
|
||||||
"sync_interval": grpo_sync_interval,
|
"sync_interval": kwargs.pop("grpo_sync_interval"),
|
||||||
}
|
}
|
||||||
|
|
||||||
executor_kwargs = {
|
executor_kwargs = {
|
||||||
@@ -391,11 +462,12 @@ def train(
|
|||||||
|
|
||||||
optimizer_fn = partial(
|
optimizer_fn = partial(
|
||||||
create_optimizer,
|
create_optimizer,
|
||||||
**{
|
lr=kwargs.pop("max_lr"),
|
||||||
"lr": max_lr,
|
weight_decay=kwargs.pop("weight_decay"),
|
||||||
"betas": (adamw_beta1, adamw_beta2),
|
momentum=kwargs.pop("muon_momentum"),
|
||||||
"weight_decay": adamw_weight_decay,
|
nesterov=kwargs.pop("muon_nesterov"),
|
||||||
},
|
ns_steps=kwargs.pop("muon_ns_steps"),
|
||||||
|
adjust_lr_fn=kwargs.pop("muon_adjust_lr"),
|
||||||
)
|
)
|
||||||
|
|
||||||
total_steps = compute_total_steps(
|
total_steps = compute_total_steps(
|
||||||
|
|||||||
@@ -0,0 +1,61 @@
|
|||||||
|
import os
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from setuptools import setup
|
||||||
|
from setuptools.command.build_ext import build_ext as _build_ext
|
||||||
|
|
||||||
|
sys.path.insert(0, str(Path(__file__).parent))
|
||||||
|
os.makedirs("astrai/extension", exist_ok=True)
|
||||||
|
|
||||||
|
|
||||||
|
def _should_build():
|
||||||
|
force = os.environ.get("CSRC_KERNELS", "").strip().lower()
|
||||||
|
if force == "true":
|
||||||
|
return True
|
||||||
|
if force == "false":
|
||||||
|
return False
|
||||||
|
try:
|
||||||
|
import shutil
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
return shutil.which("nvcc") is not None and torch.cuda.is_available()
|
||||||
|
except Exception:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
ext_modules = []
|
||||||
|
cmdclass = {}
|
||||||
|
|
||||||
|
if _should_build():
|
||||||
|
import torch
|
||||||
|
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||||
|
|
||||||
|
from csrc.build import REGISTRY
|
||||||
|
|
||||||
|
_torch_lib = torch.utils.cpp_extension.library_paths()[0]
|
||||||
|
|
||||||
|
for name, info in REGISTRY.items():
|
||||||
|
ext_modules.append(
|
||||||
|
CUDAExtension(
|
||||||
|
f"astrai.extension.{name}",
|
||||||
|
info["sources"],
|
||||||
|
extra_compile_args={
|
||||||
|
"cxx": info["cxx_flags"],
|
||||||
|
"nvcc": info["nvcc_flags"],
|
||||||
|
},
|
||||||
|
extra_link_args=[f"-Wl,-rpath,{_torch_lib}"],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
cmdclass["build_ext"] = BuildExtension
|
||||||
|
|
||||||
|
if not cmdclass:
|
||||||
|
|
||||||
|
class _NullBuildExt(_build_ext):
|
||||||
|
def build_extensions(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
cmdclass["build_ext"] = _NullBuildExt
|
||||||
|
|
||||||
|
setup(ext_modules=ext_modules, cmdclass=cmdclass)
|
||||||
+15
-1
@@ -10,7 +10,11 @@ from astrai.config.preprocess_config import (
|
|||||||
PipelineConfig,
|
PipelineConfig,
|
||||||
ProcessingConfig,
|
ProcessingConfig,
|
||||||
)
|
)
|
||||||
from astrai.preprocessing.builder import SectionedMaskBuilder
|
from astrai.preprocessing.builder import (
|
||||||
|
MultiOutputMaskBuilder,
|
||||||
|
SectionedMaskBuilder,
|
||||||
|
SingleOutputMaskBuilder,
|
||||||
|
)
|
||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
_SPECIAL_TOKENS_CONFIG = {
|
_SPECIAL_TOKENS_CONFIG = {
|
||||||
@@ -210,6 +214,16 @@ def builder():
|
|||||||
return SectionedMaskBuilder()
|
return SectionedMaskBuilder()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def single_builder():
|
||||||
|
return SingleOutputMaskBuilder()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def multi_builder():
|
||||||
|
return MultiOutputMaskBuilder()
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def tokenizer_dir(temp_dir, test_tokenizer):
|
def tokenizer_dir(temp_dir, test_tokenizer):
|
||||||
d = os.path.join(temp_dir, "tok")
|
d = os.path.join(temp_dir, "tok")
|
||||||
|
|||||||
@@ -8,7 +8,9 @@ from astrai.config.preprocess_config import (
|
|||||||
)
|
)
|
||||||
from astrai.preprocessing.builder import (
|
from astrai.preprocessing.builder import (
|
||||||
MaskBuilderFactory,
|
MaskBuilderFactory,
|
||||||
|
MultiOutputMaskBuilder,
|
||||||
SectionedMaskBuilder,
|
SectionedMaskBuilder,
|
||||||
|
SingleOutputMaskBuilder,
|
||||||
)
|
)
|
||||||
from tests.data.conftest import (
|
from tests.data.conftest import (
|
||||||
_CHAT_SECTIONS,
|
_CHAT_SECTIONS,
|
||||||
@@ -272,12 +274,18 @@ def test_sectioned_text_too_short(test_tokenizer, builder):
|
|||||||
|
|
||||||
def test_factory_registered():
|
def test_factory_registered():
|
||||||
names = MaskBuilderFactory.list_registered()
|
names = MaskBuilderFactory.list_registered()
|
||||||
|
assert "single" in names
|
||||||
|
assert "multi" in names
|
||||||
assert "sectioned" in names
|
assert "sectioned" in names
|
||||||
|
|
||||||
|
|
||||||
def test_factory_create():
|
def test_factory_create():
|
||||||
builder_obj = MaskBuilderFactory.create("sectioned")
|
single = MaskBuilderFactory.create("single")
|
||||||
assert isinstance(builder_obj, SectionedMaskBuilder)
|
assert isinstance(single, SingleOutputMaskBuilder)
|
||||||
|
multi = MaskBuilderFactory.create("multi")
|
||||||
|
assert isinstance(multi, MultiOutputMaskBuilder)
|
||||||
|
sectioned = MaskBuilderFactory.create("sectioned")
|
||||||
|
assert isinstance(sectioned, SectionedMaskBuilder)
|
||||||
|
|
||||||
|
|
||||||
def test_dpo_chat_basic(chat_tokenizer, builder):
|
def test_dpo_chat_basic(chat_tokenizer, builder):
|
||||||
@@ -367,3 +375,59 @@ def test_grpo_single_reward(chat_tokenizer, builder):
|
|||||||
}
|
}
|
||||||
result = builder.build(item, config, chat_tokenizer)
|
result = builder.build(item, config, chat_tokenizer)
|
||||||
assert result["rewards"] == [0.9]
|
assert result["rewards"] == [0.9]
|
||||||
|
|
||||||
|
|
||||||
|
def test_single_builder_matches_facade(chat_tokenizer, builder, single_builder):
|
||||||
|
config = make_chat_config()
|
||||||
|
item = {
|
||||||
|
"messages": [
|
||||||
|
{"role": "user", "content": "What is 2+2?"},
|
||||||
|
{"role": "assistant", "content": "4"},
|
||||||
|
]
|
||||||
|
}
|
||||||
|
facade_result = builder.build(item, config, chat_tokenizer)
|
||||||
|
single_result = single_builder.build(item, config, chat_tokenizer)
|
||||||
|
assert single_result == facade_result
|
||||||
|
|
||||||
|
|
||||||
|
def test_single_builder_rejects_multi_config(chat_tokenizer, single_builder):
|
||||||
|
config = make_dpo_chat_config()
|
||||||
|
item = {
|
||||||
|
"chosen": [
|
||||||
|
{"role": "user", "content": "What is 2+2?"},
|
||||||
|
{"role": "assistant", "content": "4"},
|
||||||
|
],
|
||||||
|
"rejected": [
|
||||||
|
{"role": "user", "content": "What is 2+2?"},
|
||||||
|
{"role": "assistant", "content": "5"},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
assert single_builder.build(item, config, chat_tokenizer) is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_multi_builder_matches_facade(chat_tokenizer, builder, multi_builder):
|
||||||
|
config = make_dpo_chat_config()
|
||||||
|
item = {
|
||||||
|
"chosen": [
|
||||||
|
{"role": "user", "content": "What is 2+2?"},
|
||||||
|
{"role": "assistant", "content": "4"},
|
||||||
|
],
|
||||||
|
"rejected": [
|
||||||
|
{"role": "user", "content": "What is 2+2?"},
|
||||||
|
{"role": "assistant", "content": "5"},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
facade_result = builder.build(item, config, chat_tokenizer)
|
||||||
|
multi_result = multi_builder.build(item, config, chat_tokenizer)
|
||||||
|
assert multi_result == facade_result
|
||||||
|
|
||||||
|
|
||||||
|
def test_multi_builder_rejects_single_config(chat_tokenizer, multi_builder):
|
||||||
|
config = make_chat_config()
|
||||||
|
item = {
|
||||||
|
"messages": [
|
||||||
|
{"role": "user", "content": "What is 2+2?"},
|
||||||
|
{"role": "assistant", "content": "4"},
|
||||||
|
]
|
||||||
|
}
|
||||||
|
assert multi_builder.build(item, config, chat_tokenizer) is None
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ import torch
|
|||||||
|
|
||||||
from astrai.inference import (
|
from astrai.inference import (
|
||||||
Allocator,
|
Allocator,
|
||||||
KVCache,
|
PageCache,
|
||||||
PagePool,
|
PagePool,
|
||||||
PrefixCache,
|
PrefixCache,
|
||||||
Storage,
|
Storage,
|
||||||
@@ -161,7 +161,7 @@ def test_task_table_pop():
|
|||||||
|
|
||||||
|
|
||||||
def test_kv_cache_task_extend_allocates():
|
def test_kv_cache_task_extend_allocates():
|
||||||
cache = KVCache(
|
cache = PageCache(
|
||||||
n_layers=1,
|
n_layers=1,
|
||||||
n_pages=8,
|
n_pages=8,
|
||||||
page_size=64,
|
page_size=64,
|
||||||
@@ -177,7 +177,7 @@ def test_kv_cache_task_extend_allocates():
|
|||||||
|
|
||||||
|
|
||||||
def test_kv_cache_task_extend_fails_when_pool_full():
|
def test_kv_cache_task_extend_fails_when_pool_full():
|
||||||
cache = KVCache(
|
cache = PageCache(
|
||||||
n_layers=1,
|
n_layers=1,
|
||||||
n_pages=2,
|
n_pages=2,
|
||||||
page_size=64,
|
page_size=64,
|
||||||
|
|||||||
Reference in New Issue
Block a user