Compare commits
120
Commits
8ab7564d02
...
v1.3.10
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
31d33ccdf0 | ||
|
|
88ec786e39 | ||
|
|
663ef900fc | ||
|
|
7d478a54db | ||
|
|
f3eaaef842 | ||
|
|
d655b65027 | ||
|
|
31c22dc043 | ||
|
|
17127f8b3c | ||
|
|
d7695b40e3 | ||
|
|
fc62890e70 | ||
|
|
f433672140 | ||
|
|
7e1e5b6e6a | ||
|
|
553a42702d | ||
|
|
b133fc9c07 | ||
|
|
b33250dc28 | ||
|
|
a74e5b91a3 | ||
|
|
28886e4241 | ||
|
|
9d3ccfdffc | ||
|
|
a24a7b4da5 | ||
|
|
f7df02f9a3 | ||
|
|
ee450686f3 | ||
|
|
2565755e45 | ||
|
|
d08a92c7bd | ||
|
|
a1ea26d367 | ||
|
|
c17aa0dc54 | ||
|
|
b12b24eadc | ||
|
|
cd14d53707 | ||
|
|
e220413035 | ||
|
|
84ed2327f5 | ||
|
|
b14f301730 | ||
|
|
0654b4b916 | ||
|
|
1f0be382ad | ||
|
|
bb175fda91 | ||
|
|
13998da15a | ||
|
|
57729fd92d | ||
|
|
2c7a71a9c0 | ||
|
|
3e0007fc91 | ||
|
|
b092316385 | ||
|
|
9bcd696580 | ||
|
|
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 | ||
|
|
8999ca89b8 | ||
|
|
1adca39cd8 | ||
|
|
204873fa2f | ||
|
|
a5c1de6b1b | ||
|
|
27524ad085 | ||
|
|
27d1921d9c | ||
|
|
70c0e5de90 | ||
|
|
dfb151537b | ||
|
|
500c605fad | ||
|
|
dc9faca3b1 | ||
|
|
aabb0d83e9 | ||
|
|
44579ea6dc | ||
|
|
0f1fcb079f | ||
|
|
84d4769163 | ||
|
|
bf09a35c95 | ||
|
|
6715461a36 | ||
|
|
b4587c5d08 | ||
|
|
88ec63121d | ||
|
|
01d2da2893 | ||
|
|
25d4ea3f91 | ||
|
|
39985840c7 | ||
|
|
b1adc40cfb | ||
|
|
7348bac6ab |
@@ -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
|
||||
+16
-3
@@ -5,8 +5,16 @@
|
||||
!*/
|
||||
|
||||
# Allow specific file types and root files
|
||||
!*.py
|
||||
!*.sh
|
||||
!astrai/**/*.py
|
||||
!scripts/**/*.py
|
||||
!tests/**/*.py
|
||||
!csrc/**/*.py
|
||||
|
||||
!csrc/**/*.cu
|
||||
!csrc/**/*.h
|
||||
!csrc/**/*.cuh
|
||||
|
||||
!scripts/**/*.sh
|
||||
|
||||
# Allow GitHub files
|
||||
!/.github/**
|
||||
@@ -20,4 +28,9 @@
|
||||
!/CONTRIBUTING.md
|
||||
!/LICENSE
|
||||
!/pyproject.toml
|
||||
!/README.md
|
||||
!/README.md
|
||||
# Allow extension modules (only source .py)
|
||||
!/astrai/extension/**/*.py
|
||||
|
||||
# Allow build files
|
||||
!/setup.py
|
||||
|
||||
+1
-1
@@ -23,7 +23,7 @@ COPY astrai/ ./astrai/
|
||||
COPY pyproject.toml .
|
||||
RUN pip install --no-cache-dir --upgrade pip \
|
||||
&& pip install --no-cache-dir . \
|
||||
--extra-index-url https://download.pytorch.org/whl/cu126
|
||||
--extra-index-url https://download.pytorch.org/whl/cu128
|
||||
|
||||
# Production stage
|
||||
FROM ubuntu:24.04 AS production
|
||||
|
||||
@@ -9,7 +9,7 @@
|
||||
<div align="center">
|
||||
<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/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/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks">
|
||||
</div>
|
||||
@@ -20,7 +20,7 @@
|
||||
<a href="assets/docs/README-zh-CN.md">中文</a> •
|
||||
<a href="https://github.com/ViperEkura/AstrAI/issues">Issue Tracker</a> •
|
||||
<a href="https://github.com/ViperEkura/AstrAI/discussions">Discussions</a> •
|
||||
<a href="https://huggingface.co/ViperEk/">HuggingFace</a>
|
||||
<a href="https://huggingface.co/ViperEkura">HuggingFace</a>
|
||||
</div>
|
||||
|
||||
<br>
|
||||
@@ -59,8 +59,9 @@ End-to-end walkthrough in 5 steps:
|
||||
```bash
|
||||
git clone https://github.com/ViperEkura/AstrAI.git
|
||||
cd AstrAI
|
||||
pip install -e .
|
||||
# pip install -e ".[dev]" # optional: dev dependencies (pytest, ruff)
|
||||
pip install -e . # pure PyTorch (no CUDA kernels)
|
||||
# CSRC_KERNELS=true pip install -e . --no-build-isolation # optional: fused CUDA kernels
|
||||
# pip install -e ".[dev]" # dev dependencies (pytest, ruff)
|
||||
```
|
||||
|
||||
**2. Download model**
|
||||
@@ -102,9 +103,7 @@ nohup python scripts/tools/train.py \
|
||||
--warmup_ratio=0.05 \
|
||||
--max_lr=1e-4 \
|
||||
--max_grad_norm=1.0 \
|
||||
--adamw_beta1=0.9 \
|
||||
--adamw_beta2=0.95 \
|
||||
--adamw_weight_decay=0.01 \
|
||||
--weight_decay=0.1 \
|
||||
--window_size=2048 \
|
||||
--ckpt_interval=10000 \
|
||||
--ckpt_dir=./checkpoint \
|
||||
@@ -242,7 +241,7 @@ For major changes, please open an issue first to discuss what you would like to
|
||||
|
||||
- **GitHub Issues**: [Issue Tracker](https://github.com/ViperEkura/AstrAI/issues)
|
||||
- **Discussions**: [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions)
|
||||
- **HuggingFace**: [Model Hub](https://huggingface.co/ViperEk)
|
||||
- **HuggingFace**: [Model Hub](https://huggingface.co/ViperEkura)
|
||||
|
||||
### License
|
||||
|
||||
|
||||
@@ -15,7 +15,7 @@
|
||||
<div align="center">
|
||||
<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/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/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks">
|
||||
</div>
|
||||
@@ -27,7 +27,7 @@
|
||||
<a href="#chinese">中文</a> •
|
||||
<a href="https://github.com/ViperEkura/AstrAI/issues">问题追踪</a> •
|
||||
<a href="https://github.com/ViperEkura/AstrAI/discussions">讨论区</a> •
|
||||
<a href="https://huggingface.co/ViperEk">HuggingFace</a>
|
||||
<a href="https://huggingface.co/ViperEkura">HuggingFace</a>
|
||||
</div>
|
||||
<br>
|
||||
|
||||
@@ -65,8 +65,9 @@
|
||||
```bash
|
||||
git clone https://github.com/ViperEkura/AstrAI.git
|
||||
cd AstrAI
|
||||
pip install -e .
|
||||
# pip install -e ".[dev]" # 可选:开发依赖(pytest, ruff)
|
||||
pip install -e . # 纯 PyTorch(不含 CUDA 内核)
|
||||
# CSRC_KERNELS=true pip install -e . --no-build-isolation # 可选:融合 CUDA 内核加速
|
||||
# pip install -e ".[dev]" # 可选:开发依赖(pytest, ruff)
|
||||
```
|
||||
|
||||
**2. 下载模型**
|
||||
@@ -108,9 +109,7 @@ nohup python scripts/tools/train.py \
|
||||
--warmup_ratio=0.05 \
|
||||
--max_lr=1e-4 \
|
||||
--max_grad_norm=1.0 \
|
||||
--adamw_beta1=0.9 \
|
||||
--adamw_beta2=0.95 \
|
||||
--adamw_weight_decay=0.01 \
|
||||
--weight_decay=0.1 \
|
||||
--window_size=2048 \
|
||||
--ckpt_interval=10000 \
|
||||
--ckpt_dir=./checkpoint \
|
||||
@@ -248,7 +247,7 @@ SSE 流式格式、错误码和统计端点详见[推理文档](./inference.md)
|
||||
|
||||
- **GitHub Issues**: [问题追踪](https://github.com/ViperEkura/AstrAI/issues)
|
||||
- **Discussions**: [GitHub 讨论区](https://github.com/ViperEkura/AstrAI/discussions)
|
||||
- **HuggingFace**: [模型中心](https://huggingface.co/ViperEk)
|
||||
- **HuggingFace**: [模型中心](https://huggingface.co/ViperEkura)
|
||||
|
||||
### 许可证
|
||||
|
||||
|
||||
+170
-83
@@ -21,6 +21,7 @@ classDiagram
|
||||
|
||||
class BaseModelConfig {
|
||||
+Optional[str] model_type
|
||||
+float neftune_alpha
|
||||
+from_file(config_path) Self
|
||||
+to_file(config_path)
|
||||
}
|
||||
@@ -58,10 +59,11 @@ classDiagram
|
||||
+Optional[int] dim_ffn
|
||||
+Optional[int] max_len
|
||||
+Optional[float] rope_theta
|
||||
+str attn_type
|
||||
+Optional[int] n_heads
|
||||
+Optional[int] n_kv_heads
|
||||
+Optional[bool] use_qk_norm
|
||||
+Optional[bool] use_gated_attention
|
||||
+str ffn_type
|
||||
+Optional[dict] rope_scaling
|
||||
+Optional[str] pooling_type
|
||||
+Optional[bool] normalize_embeddings
|
||||
@@ -115,14 +117,13 @@ classDiagram
|
||||
+int n_epoch
|
||||
+int batch_per_device
|
||||
+int grad_accum_steps
|
||||
+float max_grad_norm
|
||||
+Optional[float] max_grad_norm
|
||||
+list gradient_checkpointing_modules
|
||||
+int start_epoch
|
||||
+int start_batch
|
||||
+int start_samples
|
||||
+str ckpt_dir
|
||||
+int ckpt_interval
|
||||
+str log_dir
|
||||
+int log_interval
|
||||
+List[str] metrics
|
||||
+Optional[LoRAConfig] lora
|
||||
+int random_seed
|
||||
@@ -136,7 +137,9 @@ classDiagram
|
||||
+str start_method
|
||||
+str device_type
|
||||
+Optional[Dataset] val_dataset
|
||||
+Optional[float] val_split
|
||||
+int val_step
|
||||
+float neftune_alpha
|
||||
+str parallel_mode
|
||||
+dict executor_kwargs
|
||||
+dict extra_kwargs
|
||||
@@ -163,6 +166,13 @@ classDiagram
|
||||
+__getitem__(index) Dict
|
||||
}
|
||||
|
||||
class RecordDataset {
|
||||
+Optional[Callable] processor
|
||||
+load(load_path, storage_type)
|
||||
+__getitem__(index)
|
||||
+__len__()
|
||||
}
|
||||
|
||||
class DPODataset {
|
||||
+__getitem__(index) Dict
|
||||
}
|
||||
@@ -174,13 +184,26 @@ classDiagram
|
||||
class Store {
|
||||
+Dict[str, List[Tensor]] _data
|
||||
+Dict[str, List[int]] _cum
|
||||
+Dict[str, List[int]] _offsets
|
||||
+int _length
|
||||
+int _num_records
|
||||
+keys (property)
|
||||
+load(path)
|
||||
+fetch(begin, end, keys)
|
||||
+__len__()
|
||||
-_fetch_key(key, begin, end) Tensor
|
||||
-_normalize(raw)
|
||||
-_normalize(raw, offsets)
|
||||
}
|
||||
|
||||
class Streamable {
|
||||
<<mixin>>
|
||||
+fetch(begin, end, keys)
|
||||
-_fetch_stream_key(key, begin, end) Tensor
|
||||
}
|
||||
|
||||
class Recordable {
|
||||
<<mixin>>
|
||||
+num_records (property)
|
||||
+fetch_record(index, keys)
|
||||
-_fetch_record_key(key, index) Tensor
|
||||
}
|
||||
|
||||
class H5Store {
|
||||
@@ -192,6 +215,13 @@ classDiagram
|
||||
+load(path)
|
||||
}
|
||||
|
||||
class JsonlStore {
|
||||
+JsonlSource _source
|
||||
+Callable _processor
|
||||
+load(path, transform, processor)
|
||||
+fetch_record(index, keys)
|
||||
}
|
||||
|
||||
class ResumableDistributedSampler {
|
||||
+int epoch
|
||||
+int iter
|
||||
@@ -207,7 +237,7 @@ classDiagram
|
||||
+Dict _entries
|
||||
+register(name) decorator
|
||||
+create(train_type, window_size, stride) BaseDataset
|
||||
+load(train_type, load_path, window_size, stride, storage_type) BaseDataset
|
||||
+load(train_type, load_path, window_size, stride, storage_type, tokenizer_path, max_len, store) BaseDataset
|
||||
}
|
||||
}
|
||||
|
||||
@@ -215,12 +245,13 @@ classDiagram
|
||||
class Checkpoint {
|
||||
+dict state_dict
|
||||
+int epoch
|
||||
+int iteration
|
||||
+int consumed_samples
|
||||
+dict extra
|
||||
+dict meta
|
||||
+dict config
|
||||
+save(save_dir)
|
||||
+load(save_dir, broadcast) Checkpoint
|
||||
+load_any(save_dir, broadcast) Optional[Checkpoint]
|
||||
}
|
||||
}
|
||||
|
||||
@@ -350,7 +381,9 @@ classDiagram
|
||||
|
||||
class Embedding {
|
||||
+Parameter weight
|
||||
+float neftune_noise_alpha
|
||||
+forward(x) Tensor
|
||||
+set_neftune_alpha(alpha)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -372,6 +405,7 @@ classDiagram
|
||||
+List[str] paths
|
||||
+str output_dir
|
||||
+str tokenizer_path
|
||||
+AutoTokenizer tokenizer
|
||||
+BaseMaskBuilder mask_builder
|
||||
+PackingStrategy _packer
|
||||
+PositionIdStrategy _position_id
|
||||
@@ -379,6 +413,18 @@ classDiagram
|
||||
+transform(item) Optional[dict]
|
||||
+run()
|
||||
+_flush(domains, shard_idx)
|
||||
+_inject_doc_reset_position_ids(keys, mode, seqs) Dict
|
||||
+_inject_continuous_position_ids(tensors, mode, seqs) Dict
|
||||
+_to_tensors(keys) Dict
|
||||
}
|
||||
|
||||
class TokenizeTransform {
|
||||
+PipelineConfig config
|
||||
+AutoTokenizer tokenizer
|
||||
+BaseMaskBuilder mask_builder
|
||||
+PositionIdStrategy position_strategy
|
||||
+from_config_file(path) TokenizeTransform
|
||||
+apply(records) Dict[str, list]
|
||||
}
|
||||
}
|
||||
|
||||
@@ -407,7 +453,9 @@ classDiagram
|
||||
+Dict _entries
|
||||
+register(name) decorator
|
||||
+create(name, *args, **kwargs) T
|
||||
+get_component_class(name) Type
|
||||
+list_registered() list
|
||||
+is_registered(name) bool
|
||||
}
|
||||
|
||||
class MaskBuilderFactory {
|
||||
@@ -436,13 +484,15 @@ classDiagram
|
||||
+dict model_config
|
||||
+BaseExecutor executor
|
||||
+int epoch
|
||||
+int iteration
|
||||
+int consumed_samples
|
||||
+float loss
|
||||
+float grad_norm
|
||||
+DataLoader val_dataloader
|
||||
+float val_loss
|
||||
+int world_size
|
||||
+int rank
|
||||
+dict kwargs
|
||||
+optimizer_step() int
|
||||
}
|
||||
|
||||
class TrainContextBuilder {
|
||||
@@ -485,14 +535,13 @@ classDiagram
|
||||
}
|
||||
|
||||
class GRPOStrategy {
|
||||
+nn.Module old_model
|
||||
+nn.Module ref_model
|
||||
+float clip_eps
|
||||
+float kl_coef
|
||||
+int group_size
|
||||
+str reduction
|
||||
+int sync_interval
|
||||
+compute_loss(batch) Tensor
|
||||
+sync_ref_model()
|
||||
+sync_old_model()
|
||||
}
|
||||
|
||||
class BaseScheduler {
|
||||
@@ -542,12 +591,12 @@ classDiagram
|
||||
}
|
||||
|
||||
class GradientClippingCallback {
|
||||
+float max_grad_norm
|
||||
+Optional[float] max_grad_norm
|
||||
+on_optimizer_step(context)
|
||||
}
|
||||
|
||||
class GradientCheckpointingCallback {
|
||||
+tuple modules
|
||||
+Optional[List[type]] modules
|
||||
+on_train_begin(context)
|
||||
+on_train_end(context)
|
||||
}
|
||||
@@ -561,31 +610,29 @@ classDiagram
|
||||
+on_batch_end(context)
|
||||
+on_train_end(context)
|
||||
+on_error(context)
|
||||
+save_extra(context) dict$
|
||||
+save_extra(context) dict
|
||||
}
|
||||
|
||||
class ProgressBarCallback {
|
||||
+int num_epoch
|
||||
+int log_interval
|
||||
+IO file
|
||||
+tqdm progress_bar
|
||||
+on_epoch_begin(context)
|
||||
+on_batch_end(context)
|
||||
+on_optimizer_step(context)
|
||||
+on_epoch_end(context)
|
||||
}
|
||||
|
||||
class MetricLoggerCallback {
|
||||
class MetricCallback {
|
||||
+Path log_dir
|
||||
+int save_interval
|
||||
+int log_interval
|
||||
+List[str] metrics
|
||||
+on_batch_end(context)
|
||||
+int val_step
|
||||
+on_optimizer_step(context)
|
||||
+on_epoch_end(context)
|
||||
+on_train_end(context)
|
||||
+on_error(context)
|
||||
}
|
||||
|
||||
class ValidationCallback {
|
||||
-_run_validation(context)
|
||||
+on_optimizer_step(context)
|
||||
}
|
||||
|
||||
class CallbackFactory {
|
||||
@@ -594,18 +641,6 @@ classDiagram
|
||||
+create(name, **kwargs) TrainCallback
|
||||
}
|
||||
|
||||
class Muon {
|
||||
+float lr
|
||||
+float momentum
|
||||
+float weight_decay
|
||||
+bool nesterov
|
||||
+int ns_steps
|
||||
+Optional[float] adamw_lr
|
||||
+tuple adamw_betas
|
||||
+float adamw_eps
|
||||
+float adamw_wd
|
||||
+step(closure) Optional[float]
|
||||
}
|
||||
}
|
||||
|
||||
namespace inference {
|
||||
@@ -684,20 +719,44 @@ classDiagram
|
||||
}
|
||||
|
||||
class KVCache {
|
||||
-PagePool _pool
|
||||
-Storage _storage
|
||||
-TaskTable _table
|
||||
+int page_size
|
||||
<<abstract>>
|
||||
+task_alloc(task_id, prompt_ids) bool
|
||||
+task_free(task_id)
|
||||
+task_extend(task_id, pos) bool
|
||||
+task_cached(task_id) int
|
||||
+task_record_hashes(task_id, prompt_ids, start_logical_page)
|
||||
+make_table_tensor(task_ids, device) Tensor
|
||||
+bind(page_table, total_len) KvcacheView
|
||||
+bind_tasks(task_ids, total_len, device) CacheView
|
||||
}
|
||||
|
||||
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
|
||||
+Tensor _page_table
|
||||
+int _total_len
|
||||
@@ -705,6 +764,14 @@ classDiagram
|
||||
+gather(layer_id) Tuple[Tensor, Tensor]
|
||||
}
|
||||
|
||||
class ContiguousCacheView {
|
||||
-ContiguousCache _cache
|
||||
+Tensor _batch_indices
|
||||
+int _total_len
|
||||
+write(layer_id, k, v)
|
||||
+gather(layer_id) Tuple[Tensor, Tensor]
|
||||
}
|
||||
|
||||
class TaskTable {
|
||||
+set(task_id, page_table, cached)
|
||||
+get(task_id) List[int]
|
||||
@@ -714,23 +781,22 @@ classDiagram
|
||||
+table_tensor(task_ids, device) Tensor
|
||||
}
|
||||
|
||||
class Task {
|
||||
+str task_id
|
||||
+List prompt_ids
|
||||
+Optional[int] max_tokens
|
||||
+float temperature
|
||||
+float top_p
|
||||
+int top_k
|
||||
+TaskStatus status
|
||||
+List output_ids
|
||||
+int input_tokens
|
||||
+int output_tokens
|
||||
+float arrival_time
|
||||
+Optional[float] finish_time
|
||||
+Optional[Callable] stream_callback
|
||||
+int next_pos
|
||||
+is_finished(stop_ids) bool
|
||||
}
|
||||
class Task {
|
||||
+str task_id
|
||||
+List prompt_ids
|
||||
+Optional[int] max_tokens
|
||||
+float temperature
|
||||
+float top_p
|
||||
+int top_k
|
||||
+TaskStatus status
|
||||
+List output_ids
|
||||
+int input_tokens
|
||||
+int output_tokens
|
||||
+float arrival_time
|
||||
+Optional[float] finish_time
|
||||
+int next_pos
|
||||
+is_finished(stop_ids) bool
|
||||
}
|
||||
|
||||
class TaskStatus {
|
||||
<<enumeration>>
|
||||
@@ -810,7 +876,9 @@ classDiagram
|
||||
|
||||
class ChatMessage {
|
||||
+str role
|
||||
+str content
|
||||
+Optional[str] content
|
||||
+Optional[List[Dict]] tool_calls
|
||||
+Optional[str] tool_call_id
|
||||
}
|
||||
|
||||
class ChatCompletionRequest {
|
||||
@@ -827,6 +895,8 @@ classDiagram
|
||||
+Optional[float] frequency_penalty
|
||||
+Optional[Dict[int, float]] logit_bias
|
||||
+Optional[str] user
|
||||
+Optional[List[ToolDef]] tools
|
||||
+Optional[Union[str, Dict]] tool_choice
|
||||
}
|
||||
|
||||
class AnthropicMessage {
|
||||
@@ -850,7 +920,7 @@ classDiagram
|
||||
<<abstract>>
|
||||
+prepare(request, engine) Tuple[str, GenContext, List[str]]
|
||||
+format_stream_start(ctx) List[str]
|
||||
+format_chunk(token) str
|
||||
+format_chunk(token) List[str]
|
||||
+format_stream_end(ctx, stop) List[str]
|
||||
+format_response(ctx, content, stop) Dict
|
||||
}
|
||||
@@ -858,7 +928,7 @@ classDiagram
|
||||
class OpenAIResponseBuilder {
|
||||
+prepare(request, engine) Tuple
|
||||
+format_stream_start(ctx) List[str]
|
||||
+format_chunk(token) str
|
||||
+format_chunk(token) List[str]
|
||||
+format_stream_end(ctx, stop) List[str]
|
||||
+format_response(ctx, content, stop) Dict
|
||||
}
|
||||
@@ -866,7 +936,7 @@ classDiagram
|
||||
class AnthropicResponseBuilder {
|
||||
+prepare(request, engine) Tuple
|
||||
+format_stream_start(ctx) List[str]
|
||||
+format_chunk(token) str
|
||||
+format_chunk(token) List[str]
|
||||
+format_stream_end(ctx, stop) List[str]
|
||||
+format_response(ctx, content, stop) Dict
|
||||
}
|
||||
@@ -1031,14 +1101,21 @@ classDiagram
|
||||
TrainCallback <|-- GradientCheckpointingCallback
|
||||
TrainCallback <|-- CheckpointCallback
|
||||
TrainCallback <|-- ProgressBarCallback
|
||||
TrainCallback <|-- MetricLoggerCallback
|
||||
TrainCallback <|-- ValidationCallback
|
||||
TrainCallback <|-- MetricCallback
|
||||
BaseDataset <|-- SEQDataset
|
||||
BaseDataset <|-- SFTDataset
|
||||
BaseDataset <|-- DPODataset
|
||||
BaseDataset <|-- GRPODataset
|
||||
BaseDataset <|-- RecordDataset
|
||||
RecordDataset <|-- DPODataset
|
||||
RecordDataset <|-- GRPODataset
|
||||
Store <|-- H5Store
|
||||
Store <|-- MmapStore
|
||||
Store <|-- JsonlStore
|
||||
H5Store --|> Streamable
|
||||
H5Store --|> Recordable
|
||||
MmapStore --|> Streamable
|
||||
MmapStore --|> Recordable
|
||||
JsonlStore --|> Streamable
|
||||
JsonlStore --|> Recordable
|
||||
BaseSamplingStrategy <|-- TemperatureStrategy
|
||||
BaseSamplingStrategy <|-- TopKStrategy
|
||||
BaseSamplingStrategy <|-- TopPStrategy
|
||||
@@ -1071,11 +1148,15 @@ classDiagram
|
||||
ResponseBuilder <|-- OpenAIResponseBuilder
|
||||
ResponseBuilder <|-- AnthropicResponseBuilder
|
||||
BaseMaskBuilder <|-- SectionedMaskBuilder
|
||||
KVCache <|-- PageCache
|
||||
KVCache <|-- ContiguousCache
|
||||
CacheView <|-- PageCacheView
|
||||
CacheView <|-- ContiguousCacheView
|
||||
|
||||
%% --- Composition (strong ownership, part destroyed with whole) ---
|
||||
KVCache *-- PagePool
|
||||
KVCache *-- Storage
|
||||
KVCache *-- TaskTable
|
||||
PageCache *-- PagePool
|
||||
PageCache *-- Storage
|
||||
PageCache *-- TaskTable
|
||||
InferenceEngine *-- InferenceScheduler
|
||||
InferenceScheduler *-- KVCache
|
||||
InferenceScheduler *-- Executor
|
||||
@@ -1103,11 +1184,15 @@ classDiagram
|
||||
TrainContext o-- BaseScheduler
|
||||
TrainContext o-- Checkpoint
|
||||
TrainContext o-- BaseExecutor
|
||||
KvcacheView o-- Storage
|
||||
PageCacheView o-- Storage
|
||||
ContiguousCacheView o-- ContiguousCache
|
||||
SamplingPipeline o-- BaseSamplingStrategy
|
||||
BaseDataset o-- Store
|
||||
Pipeline o-- PipelineConfig
|
||||
Pipeline o-- BaseMaskBuilder
|
||||
Pipeline o-- AutoTokenizer
|
||||
TokenizeTransform o-- AutoTokenizer
|
||||
TokenizeTransform o-- BaseMaskBuilder
|
||||
|
||||
%% --- Dependency (uses temporarily) ---
|
||||
TrainConfig ..> BaseStrategy : selects
|
||||
@@ -1125,6 +1210,7 @@ classDiagram
|
||||
DecoderBlock ..> FFNFactory : uses
|
||||
StoreFactory ..> H5Store : creates
|
||||
StoreFactory ..> MmapStore : creates
|
||||
StoreFactory ..> JsonlStore : creates
|
||||
ConfigFactory ..> AutoRegressiveLMConfig : creates
|
||||
ConfigFactory ..> EncoderConfig : creates
|
||||
ExecutorFactory ..> NoneExecutor : creates
|
||||
@@ -1138,7 +1224,8 @@ classDiagram
|
||||
TrainContextBuilder ..> ResumableDistributedSampler : creates
|
||||
Checkpoint ..> Checkpoint : serializes
|
||||
CheckpointCallback ..> Checkpoint : creates
|
||||
KVCache ..> KvcacheView : binds
|
||||
PageCache ..> PageCacheView : binds
|
||||
ContiguousCache ..> ContiguousCacheView : binds
|
||||
InferenceEngine ..> GenerationRequest : uses
|
||||
InferenceEngine ..> GenerateResult : creates
|
||||
OpenAIResponseBuilder ..> ChatCompletionRequest : receives
|
||||
@@ -1149,7 +1236,7 @@ classDiagram
|
||||
%% --- Association (general usage) ---
|
||||
Trainer --> TrainConfig
|
||||
DPOStrategy --> AutoModel
|
||||
GRPOStrategy --> AutoModel
|
||||
GRPOStrategy --> AutoModel : policy/old/ref
|
||||
InferenceScheduler --> Task
|
||||
InferenceScheduler --> TaskStatus
|
||||
Task --> TaskStatus
|
||||
@@ -1166,22 +1253,22 @@ classDiagram
|
||||
| Module | Components | Description |
|
||||
|--------|------------|-------------|
|
||||
| **astrai.config** | BaseConfig, BaseModelConfig, AutoRegressiveLMConfig, EncoderConfig, ConfigFactory, TrainConfig, PipelineConfig, InputConfig, ProcessingConfig, OutputConfig | Configuration management (to_dict/from_dict, to_file/from_file) |
|
||||
| **astrai.preprocessing** | 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.preprocessing** | BaseMaskBuilder, MaskBuilderFactory, SectionedMaskBuilder, SingleOutputMaskBuilder, MultiOutputMaskBuilder, Pipeline, TokenizeTransform, filter_by_length, PackingStrategy, PackingStrategyFactory, plan_bfd, PositionIdStrategy, PositionIdStrategyFactory, StoreWriter, StoreWriterFactory, core (shared helpers) | Declarative JSON-driven data preprocessing |
|
||||
| **astrai.dataset** | BaseDataset–RecordDataset–DPO/GRPODataset, SEQDataset, SFTDataset, Store, Streamable, Recordable, H5Store, MmapStore, JsonlStore, StoreFactory, ResumableDistributedSampler, DatasetFactory | Dataset loading and management |
|
||||
| **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.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
|
||||
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy–GRPOStrategy, StrategyFactory, BaseScheduler–WSDScheduler, SchedulerFactory, TrainCallback(Protocol)–ValidationCallback, CallbackFactory, Muon | 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.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy–GRPOStrategy, StrategyFactory, BaseScheduler–WSDScheduler, SchedulerFactory, TrainCallback(Protocol)–MetricCallback, CallbackFactory | Training workflow |
|
||||
| **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.factory** | Registry, BaseFactory[T] | Component registration |
|
||||
| **astrai.factory** | BaseFactory | Component registration |
|
||||
| **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers |
|
||||
|
||||
## Design Patterns
|
||||
|
||||
| Pattern | Classes | Purpose |
|
||||
|---------|---------|---------|
|
||||
| **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory` | Decorator-based component creation |
|
||||
| **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory`, `MaskBuilderFactory`, `StoreWriterFactory`, `PackingStrategyFactory`, `PositionIdStrategyFactory` | Decorator-based component creation |
|
||||
| **Registry** | `BaseFactory` | Component registration |
|
||||
| **Strategy** | `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy` | Training strategy switching |
|
||||
| **Strategy (Sampling)** | `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations |
|
||||
@@ -1191,7 +1278,7 @@ classDiagram
|
||||
| **Context** | `TrainContext` | Unified training state bag |
|
||||
| **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction |
|
||||
| **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 |
|
||||
| **AutoModel Registry** | `AutoModel`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading |
|
||||
|
||||
@@ -1203,10 +1290,10 @@ classDiagram
|
||||
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`
|
||||
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`
|
||||
9. **Scheduler**: `SchedulerFactory` creates `CosineScheduler`/`SGDRScheduler`/`WSDScheduler`
|
||||
10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops
|
||||
11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers
|
||||
|
||||
> Document Update Time: 2026-05-30
|
||||
> Document Update Time: 2026-07-19
|
||||
|
||||
+44
-23
@@ -46,53 +46,74 @@ The output `meta.json` records the storage format, key names, dtype, total token
|
||||
|
||||
### Format Detection
|
||||
|
||||
`detect_format(load_path)` inspects the directory:
|
||||
`detect_format(load_path)` inspects the path:
|
||||
|
||||
- If `*.h5` files exist → `"h5"` (HDF5 backend)
|
||||
- If `*.bin` + `meta.json` files exist → `"bin"` (memory-mapped backend)
|
||||
- If `load_path` is a file: checks suffix — `.h5`/`.hdf5` → `"h5"`, `.jsonl` → `"jsonl"`, unknown suffix raises `ValueError`
|
||||
- If `load_path` is a directory: recursively globs for `*.h5`/`*.hdf5` files → `"h5"`, `*.bin` + `**/meta.json` → `"bin"`, or `*.jsonl` + `dataset_config.json` → `"jsonl"`
|
||||
|
||||
### Store Backends
|
||||
|
||||
Storage format is auto-detected by `detect_format()`; backends are dispatched via registry:
|
||||
|
||||
```
|
||||
StoreFactory.create("h5") → H5Store
|
||||
StoreFactory.create("bin") → MmapStore
|
||||
StoreFactory.create("h5") → H5Store
|
||||
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).
|
||||
All three inherit `Store` (base, owns `_data`/`_cum`/`_offsets`/`_normalize`) plus the `Streamable` and `Recordable` mixins, so every backend supports both `fetch(begin, end, keys)` (stream) and `fetch_record(index, keys)` (record) APIs.
|
||||
|
||||
**MmapStore**: Memory-maps `.bin` files. OS page cache sharing is native — no explicit `share_memory_()` needed. Uses `torch.from_numpy(np.memmap(...))`.
|
||||
**H5Store**: Reads HDF5 files. Tensors are loaded into host memory and normalized into segmented storage. `segments_are_records=True` — each `data_i` dataset is one record.
|
||||
|
||||
Both backends normalise tensors into `Store._data[Dict[str, List[Tensor]]]` + `Store._cum[Dict[str, List[int]]]` (cumulative lengths for bisect-based indexing).
|
||||
**MmapStore**: Memory-maps `.bin` files. OS page cache sharing is native — no explicit `share_memory_()` needed. Uses `torch.from_numpy(np.memmap(...))`. `segments_are_records=False` — bin segments are contiguous streams; record access is driven by `_offsets` (written when `save_bin(..., record_keys=...)` was used at preprocessing time).
|
||||
|
||||
**JsonlStore**: On-the-fly tokenization of raw JSONL files at load time. Requires a `dataset_config.json` alongside the `.jsonl` files following the same `PipelineConfig` schema with an additional `tokenizer_path` field. Two modes: eager (default, applies `TokenizeTransform` to all records at load) and lazy (`processor=fn` given, defers tokenisation to `fetch_record` — used by DPO/GRPO).
|
||||
|
||||
All backends normalise tensors into `Store._data[Dict[str, List[Tensor]]]` + `Store._cum[Dict[str, List[int]]]` (cumulative lengths for bisect-based stream indexing) + `Store._offsets[Dict[str, List[int]]]` (per-record offsets for record-mode indexing). Nested keys (GRPO `responses`/`masks` as `List[List[Tensor]]`) are stored as-is and excluded from both bookkeepings — they are only accessed record-by-record.
|
||||
|
||||
## Data Keys by Training Type
|
||||
|
||||
| Type | Storage Keys |
|
||||
|------|-------------|
|
||||
| `seq` | `sequence` (→ input_ids, target_ids via offset-by-1) |
|
||||
| `sft` | `sequence`, `loss_mask`, `position_ids` |
|
||||
| `dpo` | `chosen`, `rejected`, `chosen_mask`, `rejected_mask` |
|
||||
| `grpo` | `prompts`, `responses`, `masks`, `rewards` |
|
||||
| Type | Storage Keys | Access Mode |
|
||||
|------|-------------|-------------|
|
||||
| `seq` | `sequence` (→ input_ids, target_ids via offset-by-1) | stream (`fetch`) |
|
||||
| `sft` | `sequence`, `loss_mask`, `position_ids` | stream (`fetch`) |
|
||||
| `dpo` | `chosen`, `rejected`, `chosen_mask`, `rejected_mask` | record (`fetch_record`) |
|
||||
| `grpo` | `prompts`, `responses`, `masks`, `rewards` | record (`fetch_record`) |
|
||||
|
||||
## Dataset Architecture
|
||||
|
||||
```
|
||||
DatasetFactory.load(train_type, load_path, window_size, stride=None, storage_type=None)
|
||||
DatasetFactory.load(train_type, load_path, window_size, stride=None,
|
||||
storage_type=None, tokenizer_path=None,
|
||||
max_len=2048, store=None)
|
||||
→ BaseDataset.load(load_path, storage_type=None)
|
||||
→ detect_format(load_path)
|
||||
→ StoreFactory.create(storage_type)
|
||||
→ Store.load(load_path)
|
||||
→ H5Store._normalize() / MmapStore._normalize()
|
||||
→ Store._data[Dict[str, List[Tensor]]] + _cum[Dict[str, List[int]]]
|
||||
→ BaseDataset.__getitem__(idx)
|
||||
→ get_index(idx) → [begin, end)
|
||||
→ Store.fetch(begin, end, keys) → Tensor / Dict[str, Tensor]
|
||||
→ _normalize(raw) # base Store, shared by both backends
|
||||
→ Store._data[Dict[str, List[Tensor]]]
|
||||
+ _cum[Dict[str, List[int]]] (stream mode)
|
||||
+ _offsets[Dict[str, List[int]]] (record mode)
|
||||
|
||||
Stream datasets (SEQ/SFT):
|
||||
BaseDataset.__getitem__(idx)
|
||||
→ get_index(idx) → [begin, end)
|
||||
→ Store.fetch(begin, end, keys) → Tensor / Dict[str, Tensor]
|
||||
|
||||
Record datasets (DPO/GRPO via RecordDataset):
|
||||
RecordDataset.__getitem__(idx)
|
||||
→ Store.fetch_record(idx, keys) → Tensor / Dict[str, Tensor]
|
||||
```
|
||||
|
||||
`window_size` = max input length, `stride` = step between consecutive samples (defaults to `window_size`, optional). `storage_type` defaults to `None` (auto-detect via `detect_format`).
|
||||
Class hierarchy: `BaseDataset` ← `SEQDataset` / `SFTDataset` (stream); `BaseDataset` ← `RecordDataset` ← `DPODataset` / `GRPODataset` (record).
|
||||
|
||||
`Store.fetch(begin, end, keys)` accepts a single key (`str`) returning a `Tensor`, or a list of keys returning `Dict[str, Tensor]`. Internally uses `bisect` across multi-segment tensors. Raises `RuntimeError("Store not loaded")` if called before `load()`.
|
||||
`window_size` = max input length, `stride` = step between consecutive samples (defaults to `window_size`, optional). Only meaningful for stream datasets — record datasets ignore both. `storage_type` defaults to `None` (auto-detect via `detect_format`).
|
||||
|
||||
`tokenizer_path` triggers lazy on-the-fly tokenisation for record datasets on raw JSONL (DPO builds a `dpo_processor`; SEQ/SFT/pre-tokenised backends ignore it). `store` (pre-built `Store`) bypasses `load_path`/`storage_type`/`tokenizer_path` entirely — the caller controls Store construction.
|
||||
|
||||
`Store.fetch(begin, end, keys)` (stream mode, on `Streamable`): accepts a single key (`str`) returning a `Tensor`, or a list of keys returning `Dict[str, Tensor]`. Internally uses `bisect` across multi-segment tensors. Raises `RuntimeError("Store not loaded")` if called before `load()`.
|
||||
|
||||
`Store.fetch_record(index, keys)` (record mode, on `Recordable`): same key API. Uses `_offsets[key]` when present (bin layout with per-record offsets), otherwise indexes `_data[key]` directly (H5/JSONL where each segment is one record).
|
||||
|
||||
## Sampler
|
||||
|
||||
@@ -106,4 +127,4 @@ DatasetFactory.load(train_type, load_path, window_size, stride=None, storage_typ
|
||||
|
||||
Standard PyTorch `DataLoader` with configurable `batch_size`, `num_workers`, `pin_memory`, `prefetch_factor`. Sampler produces indices; dataloader fetches tensor batches via `__getitem__`.
|
||||
|
||||
> Document Update Time: 2026-06-19
|
||||
> Document Update Time: 2026-07-19
|
||||
|
||||
+39
-36
@@ -23,29 +23,40 @@ RoPE is applied **before** KV cache write, not after — otherwise position enco
|
||||
|
||||
## KVCache System
|
||||
|
||||
Six classes (plus two helpers) working together:
|
||||
Seven classes working together, with two concrete cache implementations:
|
||||
|
||||
### ContiguousCache (default)
|
||||
|
||||
```
|
||||
KVCache (facade)
|
||||
├── PagePool orchestrates page allocation + prefix matching
|
||||
│ ├── 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())
|
||||
ContiguousCache (simple contiguous per-slot cache)
|
||||
├── ContiguousCacheView bundles k/v tensors + slot indices for attention layers
|
||||
```
|
||||
|
||||
`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
|
||||
|
||||
`InferenceScheduler` runs a daemon thread with a 4-phase loop:
|
||||
|
||||
```
|
||||
1. Cleanup → Remove finished tasks, free KV pages
|
||||
2. Refill → Pop from waiting_queue, task_alloc pages, activate
|
||||
1. Cleanup → Remove finished tasks, free KV cache slots/pages
|
||||
2. Refill → Pop from waiting_queue, task_alloc resources, activate
|
||||
3. Prefill → Group by (prompt_len, start_pos), run full forward
|
||||
4. Decode → Pick largest same-position group, single-token forward
|
||||
4. Decode → Run single-token forward for each same-position group
|
||||
```
|
||||
|
||||
## Sampling (Strategy Pattern)
|
||||
@@ -152,12 +163,13 @@ Supports `stop_sequences` and streaming via `event: content_block_delta`.
|
||||
data: {"id":"chatcmpl-...","object":"chat.completion.chunk","created":...,"model":"astrai",
|
||||
"choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}
|
||||
|
||||
data: {"id":"chatcmpl-...","object":"chat.completion.chunk",...,
|
||||
data: {"id":"chatcmpl-...","object":"chat.completion.chunk","created":0,"model":"astrai",
|
||||
"choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]}
|
||||
|
||||
data: {"id":"chatcmpl-...","object":"chat.completion.chunk",...,
|
||||
"choices":[{"index":0,"delta":{},"finish_reason":"stop"}],
|
||||
"usage":{"prompt_tokens":5,"completion_tokens":1,"total_tokens":6}}
|
||||
data: {"id":"chatcmpl-...","object":"chat.completion.chunk","created":...,"model":"astrai",
|
||||
"choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}
|
||||
|
||||
data: {"prompt_tokens":5,"completion_tokens":1,"total_tokens":6}
|
||||
|
||||
data: [DONE]
|
||||
```
|
||||
@@ -167,7 +179,7 @@ data: [DONE]
|
||||
```
|
||||
event: message_start
|
||||
data: {"type":"message_start","message":{"id":"msg_...","model":"astrai","role":"assistant",
|
||||
"content":[],"stop_reason":null,...}}
|
||||
"content":[],"usage":{"input_tokens":0}}}
|
||||
|
||||
event: content_block_start
|
||||
data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}
|
||||
@@ -179,7 +191,7 @@ event: content_block_stop
|
||||
data: {"type":"content_block_stop","index":0}
|
||||
|
||||
event: message_delta
|
||||
data: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{...}}
|
||||
data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},"usage":{...}}
|
||||
|
||||
event: message_stop
|
||||
data: {"type":"message_stop"}
|
||||
@@ -187,26 +199,20 @@ data: {"type":"message_stop"}
|
||||
|
||||
### Error Responses
|
||||
|
||||
All endpoints use standard HTTP status codes:
|
||||
The server returns standard HTTP status codes. Pydantic validation errors (e.g. missing required fields)
|
||||
are handled automatically by FastAPI with 422 status. The only application-level error is engine initialization:
|
||||
|
||||
| Status | Meaning |
|
||||
|--------|---------|
|
||||
| 200 | Success |
|
||||
| 400 | Invalid request (bad JSON, missing fields, validation error) |
|
||||
| 405 | Method not allowed |
|
||||
| 422 | Unprocessable entity (Pydantic validation) |
|
||||
| 500 | Internal server error (model crash, OOM, scheduler failure) |
|
||||
| 503 | Service unavailable (model not loaded, engine not ready) |
|
||||
|
||||
Error response body:
|
||||
Error response body (503):
|
||||
|
||||
```json
|
||||
{
|
||||
"error": {
|
||||
"message": "Invalid request: max_tokens must be > 0",
|
||||
"type": "invalid_request_error",
|
||||
"code": 400
|
||||
}
|
||||
"detail": "Engine not initialized"
|
||||
}
|
||||
```
|
||||
|
||||
@@ -220,16 +226,13 @@ Response:
|
||||
|
||||
```json
|
||||
{
|
||||
"active_requests": 3,
|
||||
"waiting_requests": 2,
|
||||
"total_requests": 128,
|
||||
"cache_usage": 0.45,
|
||||
"tokens_generated": 10240
|
||||
"total_tasks": 128,
|
||||
"total_tokens": 10240,
|
||||
"active_tasks": 3,
|
||||
"waiting_queue": 2
|
||||
}
|
||||
```
|
||||
|
||||
`cache_usage` is the fraction of KV cache pages currently in use (0.0–1.0).
|
||||
|
||||
## Engine API
|
||||
|
||||
```python
|
||||
@@ -246,4 +249,4 @@ async for token in engine.generate_async("Hello", ...): # -> AsyncGenerator[s
|
||||
print(token)
|
||||
```
|
||||
|
||||
> Document Update Time: 2026-06-19
|
||||
> Document Update Time: 2026-07-09
|
||||
|
||||
+26
-14
@@ -26,15 +26,19 @@
|
||||
|-----------|-------------|---------|
|
||||
| `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 |
|
||||
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
|
||||
| `--max_grad_norm` | Maximum gradient norm for clipping | 1.0 |
|
||||
| `--max_grad_norm` | Maximum gradient norm for clipping (None disables) | None |
|
||||
|
||||
### Optimizer (AdamW)
|
||||
### Optimizer (MuonMix)
|
||||
|
||||
Combined optimizer: matrix parameters via **Muon**, non-matrix via **AdamW** (`fused=True`).
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--adamw_beta1` | AdamW beta1 | 0.9 |
|
||||
| `--adamw_beta2` | AdamW beta2 | 0.95 |
|
||||
| `--adamw_weight_decay` | AdamW weight decay | 0.01 |
|
||||
| `--weight_decay` | Weight decay (applied to Muon matrix params; non-matrix use 0) | 0.1 |
|
||||
| `--muon_momentum` | Muon momentum factor | 0.95 |
|
||||
| `--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
|
||||
|
||||
@@ -53,7 +57,7 @@
|
||||
| `--ckpt_interval` | Iterations between checkpoints | 5000 |
|
||||
| `--ckpt_dir` | Checkpoint save directory | checkpoint |
|
||||
| `--start_epoch` | Resume from epoch (0 = from scratch) | 0 |
|
||||
| `--start_batch` | Resume from batch iteration | 0 |
|
||||
| `--start_samples` | Resume from sample count per rank | 0 |
|
||||
|
||||
### Validation
|
||||
|
||||
@@ -67,8 +71,7 @@
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--log_dir` | Directory for metric logs | checkpoint/logs |
|
||||
| `--log_interval` | Number of batch iterations between metric logs | 100 |
|
||||
| `--metrics` | Metrics to log (e.g. --metrics loss lr val_loss) | ["loss", "lr"] |
|
||||
| `--metrics` | Metrics to log (e.g. --metrics loss lr val_loss) | ["loss", "lr", "grad_norm"] |
|
||||
|
||||
### Gradient Checkpointing
|
||||
|
||||
@@ -100,6 +103,17 @@
|
||||
| `--grpo_sync_interval` | GRPO ref_model sync interval (steps) | 200 | `grpo` |
|
||||
| `--neftune_alpha` | NEFTune noise alpha (0=disabled, typical: 5.0) | 0.0 | `sft` |
|
||||
|
||||
### Scheduler
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--schedule_type` | LR scheduler type (`cosine`, `sgdr`, `wsd`) | cosine |
|
||||
| `--min_rate` | Minimum LR as fraction of base LR | None (scheduler default: 0.01) |
|
||||
| `--cycle_length` | SGDR first cycle length in steps | None (total_steps - warmup_steps) |
|
||||
| `--t_mult` | SGDR cycle length multiplier per restart | 2 |
|
||||
| `--stable_steps` | WSD stable plateau steps | None (required for wsd) |
|
||||
| `--decay_steps` | WSD decay steps | None (total_steps - warmup_steps - stable_steps) |
|
||||
|
||||
### Usage Example
|
||||
|
||||
```bash
|
||||
@@ -116,9 +130,7 @@ nohup python scripts/tools/train.py \
|
||||
--warmup_ratio=0.05 \
|
||||
--max_lr=1e-4 \
|
||||
--max_grad_norm=1.0 \
|
||||
--adamw_beta1=0.9 \
|
||||
--adamw_beta2=0.95 \
|
||||
--adamw_weight_decay=0.01 \
|
||||
--weight_decay=0.1 \
|
||||
--window_size=2048 \
|
||||
--ckpt_interval=10000 \
|
||||
--ckpt_dir=./checkpoint \
|
||||
@@ -161,7 +173,7 @@ See [Inference Guide](inference.md) for HTTP API documentation.
|
||||
| `--top_k` | int | `30` | Top-k filtering |
|
||||
| `--top_p` | float | `0.95` | Nucleus sampling threshold |
|
||||
| `--batch_size` | int | `1` | Batch size for generation |
|
||||
| `--max_tokens` | int | `2048` | Maximum tokens to generate |
|
||||
| `--max_tokens` | int | model config `max_len` | Maximum tokens to generate |
|
||||
|
||||
Usage:
|
||||
```bash
|
||||
@@ -178,7 +190,7 @@ python scripts/tools/generate.py \
|
||||
| `input_files` | path(s) | required | Input JSONL file(s), supports glob (`data/*.jsonl`) |
|
||||
| `--output_dir`, `-o` | path | required | Output directory for processed data |
|
||||
| `--config`, `-c` | path | required | Preprocessing pipeline config (JSON) |
|
||||
| `--num_workers` | int | `4` | Number of parallel workers |
|
||||
| `--tokenizer_path` | str | `params` | Path to tokenizer directory |
|
||||
|
||||
Usage:
|
||||
```bash
|
||||
@@ -189,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-19
|
||||
@@ -1,6 +1,6 @@
|
||||
# Preprocessing Pipeline
|
||||
|
||||
Declarative JSON-driven data preprocessing. One `SectionedMaskBuilder` handles all formats via `input.sections` (single-output) or `input.sources` (multi-output).
|
||||
Declarative JSON-driven data preprocessing. `MaskBuilderFactory` supports three registered builders: `"single"` (single-output via `input.sections`), `"multi"` (multi-output via `input.sources`), and `"sectioned"` (façade dispatching to `single` or `multi` based on config).
|
||||
|
||||
## Contents
|
||||
|
||||
@@ -26,8 +26,9 @@ A single config file captures the entire pipeline, reusable and version-controll
|
||||
|
||||
```json
|
||||
{
|
||||
"version": 1,
|
||||
"input": {}, // sections (single) or sources (multi)
|
||||
"mask": {}, // role → "train" | "mask"
|
||||
"mask": {}, // role -> "train" | "mask"
|
||||
"mask_default": "mask",
|
||||
"preprocessing": {},
|
||||
"output": {}
|
||||
@@ -220,11 +221,12 @@ Config:
|
||||
}
|
||||
```
|
||||
|
||||
Output keys: `prompts`, `responses`, `masks`, `rewards` (float32)
|
||||
Output keys: `prompts`, `prompts_mask`, `responses`, `masks`, `rewards` (float32)
|
||||
|
||||
- `action: "value"` — extract raw values from JSONL without tokenisation
|
||||
- `list_field: true` — tokenise each list element independently, then concatenate
|
||||
- `mask_key: "masks"` — rename the auto-generated mask key (default: `responses_mask`)
|
||||
- `prompts_mask` is auto-generated (all masked) and unused by GRPOStrategy
|
||||
|
||||
---
|
||||
|
||||
@@ -266,7 +268,7 @@ When `sources` is set, `sections` is ignored.
|
||||
| `storage_format` | str | `"bin"` | `"bin"` (mmap) or `"h5"` |
|
||||
| `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"}`) |
|
||||
| `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"` |
|
||||
|
||||
---
|
||||
|
||||
@@ -274,12 +276,11 @@ When `sources` is set, `sections` is ignored.
|
||||
|
||||
### Template mode (`template: true`)
|
||||
|
||||
For each message in the field's array:
|
||||
|
||||
1. Prepend BOS token (masked)
|
||||
2. Render through `chat_template` for that single message
|
||||
3. Encode rendered text
|
||||
4. Apply mask rule for the message's role
|
||||
2. For each message in the field's array:
|
||||
1. Render through `chat_template` for that single message
|
||||
2. Encode rendered text
|
||||
3. Apply mask rule for the message's role
|
||||
|
||||
### Non-template mode
|
||||
|
||||
@@ -287,7 +288,7 @@ Encode the field value as text. Mask value is 1 (train) or 0 (mask) per the sect
|
||||
|
||||
### Text config detection
|
||||
|
||||
When no section uses `template` and all sections have `action: "train"`, the builder skips mask generation entirely — all tokens are trained.
|
||||
When no section uses `template` and all sections have `action: "train"`, the builder omits `loss_mask` from the output — all tokens are trained.
|
||||
|
||||
---
|
||||
|
||||
@@ -298,13 +299,15 @@ When no section uses `template` and all sections have `action: "train"`, the bui
|
||||
```
|
||||
output/
|
||||
__default__/
|
||||
meta.json
|
||||
sequence.bin
|
||||
loss_mask.bin
|
||||
shard_0000/
|
||||
meta.json
|
||||
sequence.bin
|
||||
loss_mask.bin
|
||||
wiki/
|
||||
meta.json
|
||||
sequence.bin
|
||||
loss_mask.bin
|
||||
shard_0000/
|
||||
meta.json
|
||||
sequence.bin
|
||||
loss_mask.bin
|
||||
```
|
||||
|
||||
### Multi-Shard (`bin`)
|
||||
@@ -324,7 +327,7 @@ output/
|
||||
loss_mask.bin
|
||||
```
|
||||
|
||||
`MmapStore` discovers all shards under the domain directory via `rglob("meta.json")`.
|
||||
For `bin` format, `MmapStore` discovers all shards under the domain directory via `rglob("meta.json")`. For `h5` format, `H5Store` discovers `.h5`/`.hdf5` files via recursive glob.
|
||||
|
||||
---
|
||||
|
||||
@@ -349,7 +352,7 @@ python scripts/tools/preprocess.py data/grpo/*.jsonl -o output/grpo/ -c configs/
|
||||
from astrai.preprocessing.pipeline import Pipeline
|
||||
from astrai.config.preprocess_config import PipelineConfig
|
||||
|
||||
config = PipelineConfig.from_json("sft.json")
|
||||
config = PipelineConfig.from_file("sft.json")
|
||||
Pipeline(
|
||||
config,
|
||||
["data_part1.jsonl", "data_part2.jsonl"],
|
||||
@@ -358,4 +361,4 @@ Pipeline(
|
||||
).run()
|
||||
```
|
||||
|
||||
> Document Update Time: 2026-06-03
|
||||
> Document Update Time: 2026-07-09
|
||||
|
||||
+28
-18
@@ -58,7 +58,9 @@ on_train_begin
|
||||
context.loss = loss.item()
|
||||
stand_loss = loss / executor.grad_accum_steps
|
||||
executor.backward(stand_loss)
|
||||
context.iteration += 1
|
||||
context.consumed_samples += (
|
||||
context.config.batch_per_device * context.world_size
|
||||
)
|
||||
on_batch_end
|
||||
|
||||
if executor.sync_gradients:
|
||||
@@ -78,13 +80,13 @@ on_train_end
|
||||
| `on_train_begin` | Before training starts | `GradientCheckpointingCallback` |
|
||||
| `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` |
|
||||
| `on_batch_begin` | Every batch | — |
|
||||
| `on_optimizer_step` | Every accumulation window | `GradientClippingCallback`, `ValidationCallback` |
|
||||
| `on_batch_end` | Every batch | `CheckpointCallback`, `MetricLoggerCallback`, `ProgressBarCallback` |
|
||||
| `on_epoch_end` | End of each epoch | `ProgressBarCallback` |
|
||||
| `on_error` | On exception during training | `CheckpointCallback`, `MetricLoggerCallback` |
|
||||
| `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricLoggerCallback`, `GradientCheckpointingCallback` |
|
||||
| `on_optimizer_step` | Every accumulation window | `GradientClippingCallback`, `MetricCallback`, `ProgressBarCallback` |
|
||||
| `on_batch_end` | Every batch | `CheckpointCallback` |
|
||||
| `on_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` |
|
||||
| `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` |
|
||||
| `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricCallback`, `GradientCheckpointingCallback` |
|
||||
|
||||
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric_logger` (JSONL, rank-0), `progress_bar` (tqdm), `gradient_clipping`, `validation` (periodic validation on val_dataset).
|
||||
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm), `gradient_clipping` (always registered; computes grad norm, clips only when `max_grad_norm` is not `None`).
|
||||
|
||||
## Strategies
|
||||
|
||||
@@ -106,7 +108,7 @@ $$
|
||||
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)
|
||||
|
||||
@@ -116,21 +118,31 @@ $$
|
||||
L_{\text{DPO}} = -\mathbb{E}\left[\log\sigma\left(\beta\log\frac{\pi_\theta(y_w\mid x)}{\pi_{\text{ref}}(y_w\mid x)} - \beta\log\frac{\pi_\theta(y_l\mid x)}{\pi_{\text{ref}}(y_l\mid x)}\right)\right]
|
||||
$$
|
||||
|
||||
Parameters: `beta=0.1`, `reduction="mean"`. Keys: `chosen`, `rejected`, `chosen_mask`, `rejected_mask`.
|
||||
Parameters: `beta=0.1`, `reduction="sum"`. Keys: `chosen`, `rejected`, `chosen_mask`, `rejected_mask`.
|
||||
|
||||
### GRPO (Group Relative Policy Optimization)
|
||||
|
||||
On-policy PPO with group-normalized advantages:
|
||||
Token-level PPO with group-normalized advantages. Advantages are derived from
|
||||
scalar per-response rewards, group-normalized, and broadcast across all response
|
||||
tokens. Only response tokens contribute to the loss (prompt tokens are masked
|
||||
out):
|
||||
|
||||
$$
|
||||
\text{Advantage}_i = \frac{r_i - \mu}{\sigma + \epsilon}
|
||||
$$
|
||||
|
||||
$$
|
||||
L_{\text{GRPO}} = -\mathbb{E}\left[\min\left(\frac{\pi_\theta}{\pi_{\text{ref}}}A,\; \text{clip}\left(\frac{\pi_\theta}{\pi_{\text{ref}}}, 1-\epsilon, 1+\epsilon\right)A\right)\right] + \lambda \cdot \mathbb{E}\left[(\log\pi_\theta - \log\pi_{\text{ref}})^2\right]
|
||||
L_{\text{GRPO}} = -\mathbb{E}_t\left[\min\left(\rho_t A,\; \text{clip}\left(\rho_t, 1-\epsilon, 1+\epsilon\right)A\right)\right] + \lambda \cdot \mathbb{E}_t\left[\frac{\pi_{\text{ref}}}{\pi_\theta} - \log\frac{\pi_{\text{ref}}}{\pi_\theta} - 1\right]
|
||||
$$
|
||||
|
||||
Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`, `sync_interval=200`, `reduction="mean"`.
|
||||
where $\rho_t = \pi_\theta(a_t|s_t) / \pi_{\text{old}}(a_t|s_t)$ is the
|
||||
per-token importance sampling ratio against the behaviour policy
|
||||
(`old_model`, synced externally between data-generation rounds) and the
|
||||
expectations are over valid response tokens. The KL term regularises
|
||||
$\pi_\theta$ towards a frozen reference model (`ref_model`, typically
|
||||
the SFT checkpoint).
|
||||
|
||||
Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`. External sync of `old_model` weights via `sync_old_model()` between data-generation rounds.
|
||||
|
||||
Keys: `prompts`, `responses`, `masks`, `rewards`.
|
||||
|
||||
@@ -158,8 +170,8 @@ Callback wraps each `DecoderBlock.forward` with `torch.utils.checkpoint.checkpoi
|
||||
## Checkpoint
|
||||
|
||||
```
|
||||
Checkpoint(state_dict, epoch, iteration, extra, meta, config)
|
||||
├── save(save_dir) rank-0 only: meta.json (epoch/iteration/timestamp) + config.json (model config) + model.safetensors + optional {key}.pt (optimizer.pt, scheduler.pt)
|
||||
Checkpoint(state_dict, epoch, consumed_samples, extra, meta, config)
|
||||
├── save(save_dir) rank-0 only: meta.json (epoch/consumed_samples/timestamp) + config.json (model config) + model.safetensors + optional {key}.pt (optimizer.pt, scheduler.pt)
|
||||
└── load(save_dir, broadcast=False) loads from local disk; set broadcast=True to broadcast metadata from rank-0
|
||||
```
|
||||
|
||||
@@ -199,9 +211,7 @@ nohup python scripts/tools/train.py \
|
||||
--warmup_ratio=0.05 \
|
||||
--max_lr=1e-4 \
|
||||
--max_grad_norm=1.0 \
|
||||
--adamw_beta1=0.9 \
|
||||
--adamw_beta2=0.95 \
|
||||
--adamw_weight_decay=0.01 \
|
||||
--weight_decay=0.1 \
|
||||
--window_size=2048 \
|
||||
--ckpt_interval=10000 \
|
||||
--ckpt_dir=./checkpoint \
|
||||
@@ -212,4 +222,4 @@ nohup python scripts/tools/train.py \
|
||||
|
||||
Full parameter reference at [params.md](params.md).
|
||||
|
||||
> Document Update Time: 2026-05-30
|
||||
> Document Update Time: 2026-07-19
|
||||
|
||||
+3
-5
@@ -1,4 +1,4 @@
|
||||
__version__ = "1.3.7"
|
||||
__version__ = "1.3.10"
|
||||
__author__ = "ViperEkura"
|
||||
|
||||
from astrai.config import (
|
||||
@@ -12,7 +12,7 @@ from astrai.config import (
|
||||
from astrai.dataset import (
|
||||
BaseDataset,
|
||||
DatasetFactory,
|
||||
ResumableDistributedSampler,
|
||||
RDSampler,
|
||||
Store,
|
||||
StoreFactory,
|
||||
)
|
||||
@@ -47,7 +47,6 @@ from astrai.trainer import (
|
||||
BaseScheduler,
|
||||
BaseStrategy,
|
||||
CallbackFactory,
|
||||
Muon,
|
||||
SchedulerFactory,
|
||||
StrategyFactory,
|
||||
TrainCallback,
|
||||
@@ -75,11 +74,10 @@ __all__ = [
|
||||
"GenerationRequest",
|
||||
"InferenceEngine",
|
||||
"LoRAConfig",
|
||||
"Muon",
|
||||
"Pipeline",
|
||||
"PipelineConfig",
|
||||
"ProtocolHandler",
|
||||
"ResumableDistributedSampler",
|
||||
"RDSampler",
|
||||
"SamplingPipeline",
|
||||
"SchedulerFactory",
|
||||
"Store",
|
||||
|
||||
@@ -20,6 +20,7 @@ class BaseModelConfig(BaseConfig):
|
||||
"""Base config with ``model_type`` dispatch and file I/O."""
|
||||
|
||||
model_type: Optional[str] = None
|
||||
neftune_alpha: float = 0.0
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -70,10 +71,12 @@ class EncoderConfig(BaseModelConfig):
|
||||
rope_theta: Optional[float] = None
|
||||
rope_scaling: Optional[dict] = None
|
||||
|
||||
attn_type: str = "gqa"
|
||||
n_heads: Optional[int] = None
|
||||
n_kv_heads: Optional[int] = None
|
||||
use_qk_norm: Optional[bool] = None
|
||||
use_gated_attention: Optional[bool] = None
|
||||
|
||||
ffn_type: str = "mlp"
|
||||
pooling_type: Optional[str] = None
|
||||
normalize_embeddings: Optional[bool] = None
|
||||
|
||||
@@ -96,7 +96,7 @@ class OutputConfig(BaseConfig):
|
||||
storage_format: str = "bin"
|
||||
max_tokens_per_shard: int = 100_000_000
|
||||
dtype: Dict[str, str] = field(default_factory=dict)
|
||||
position_ids_mode: str = "none"
|
||||
position_ids_mode: str = "doc_reset"
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -37,8 +37,9 @@ class TrainConfig(BaseConfig):
|
||||
grad_accum_steps: int = field(
|
||||
default=1, metadata={"help": "Number of iterations between steps."}
|
||||
)
|
||||
max_grad_norm: float = field(
|
||||
default=1.0, metadata={"help": "Maximum gradient norm."}
|
||||
max_grad_norm: Optional[float] = field(
|
||||
default=None,
|
||||
metadata={"help": "Maximum gradient norm. None disables clipping."},
|
||||
)
|
||||
gradient_checkpointing_modules: List[str] = field(
|
||||
default_factory=list,
|
||||
@@ -47,14 +48,18 @@ class TrainConfig(BaseConfig):
|
||||
|
||||
# checkpoint setting
|
||||
start_epoch: int = field(default=0, metadata={"help": "Start epoch for training."})
|
||||
start_batch: int = field(
|
||||
default=0, metadata={"help": "Start batch iteration for training."}
|
||||
start_samples: int = field(
|
||||
default=0,
|
||||
metadata={
|
||||
"help": "Start samples count (per rank). Superseded by checkpoint consumed_samples."
|
||||
},
|
||||
)
|
||||
ckpt_dir: str = field(
|
||||
default="./checkpoint", metadata={"help": "Checkpoint directory."}
|
||||
)
|
||||
ckpt_interval: int = field(
|
||||
default=5000, metadata={"help": "Number of iterations between checkpoints."}
|
||||
default=5000,
|
||||
metadata={"help": "Number of optimizer steps between checkpoints."},
|
||||
)
|
||||
|
||||
# lora setting
|
||||
@@ -67,12 +72,8 @@ class TrainConfig(BaseConfig):
|
||||
log_dir: str = field(
|
||||
default="./checkpoint/logs", metadata={"help": "Directory for metric logs."}
|
||||
)
|
||||
log_interval: int = field(
|
||||
default=100,
|
||||
metadata={"help": "Number of batch iterations between metric logs."},
|
||||
)
|
||||
metrics: List[str] = field(
|
||||
default_factory=lambda: ["loss", "lr"],
|
||||
default_factory=lambda: ["loss", "lr", "grad_norm"],
|
||||
metadata={"help": "Metrics to record during training."},
|
||||
)
|
||||
|
||||
@@ -87,6 +88,10 @@ class TrainConfig(BaseConfig):
|
||||
pin_memory: bool = field(
|
||||
default=False, metadata={"help": "Pin memory for dataloader."}
|
||||
)
|
||||
collate_fn: Optional[Callable[[List[Any]], Any]] = field(
|
||||
default=None,
|
||||
metadata={"help": "Collate function for dataloader (e.g. dpo_collate_fn)."},
|
||||
)
|
||||
|
||||
# distributed training
|
||||
nprocs: int = field(
|
||||
|
||||
@@ -1,14 +1,21 @@
|
||||
from astrai.dataset.dataset import (
|
||||
BaseDataset,
|
||||
DatasetFactory,
|
||||
dpo_collate_fn,
|
||||
grpo_collate_fn,
|
||||
)
|
||||
from astrai.dataset.sampler import ResumableDistributedSampler
|
||||
from astrai.dataset.sampler import RDSampler
|
||||
from astrai.dataset.storage import (
|
||||
H5Store,
|
||||
JsonlStore,
|
||||
MmapStore,
|
||||
Recordable,
|
||||
Store,
|
||||
StoreFactory,
|
||||
Streamable,
|
||||
detect_format,
|
||||
)
|
||||
from astrai.serialization import (
|
||||
load_bin,
|
||||
load_h5,
|
||||
save_bin,
|
||||
@@ -18,14 +25,19 @@ from astrai.dataset.storage import (
|
||||
__all__ = [
|
||||
"BaseDataset",
|
||||
"DatasetFactory",
|
||||
"dpo_collate_fn",
|
||||
"grpo_collate_fn",
|
||||
"Store",
|
||||
"Streamable",
|
||||
"Recordable",
|
||||
"StoreFactory",
|
||||
"H5Store",
|
||||
"MmapStore",
|
||||
"JsonlStore",
|
||||
"detect_format",
|
||||
"save_h5",
|
||||
"load_h5",
|
||||
"save_bin",
|
||||
"load_bin",
|
||||
"ResumableDistributedSampler",
|
||||
"RDSampler",
|
||||
]
|
||||
|
||||
+401
-180
@@ -1,7 +1,31 @@
|
||||
"""Dataset implementations with factory pattern for training."""
|
||||
"""Dataset implementations for training.
|
||||
|
||||
Composition over inheritance — every dataset is a thin wrapper that
|
||||
binds a :class:`Store` to a particular train-type's key mapping. All
|
||||
sample-id → token/record indexing lives on the Store; datasets never
|
||||
know about window/stride math or segment layouts.
|
||||
|
||||
Class hierarchy:
|
||||
|
||||
BaseDataset (ABC) — holds a Store, exposes __len__/keys,
|
||||
overrides __getitem__
|
||||
├── SEQDataset — next-token prediction (stream)
|
||||
├── SFTDataset — loss-mask + position_ids (stream)
|
||||
├── DPODataset — chosen/rejected pairs (record)
|
||||
└── GRPODataset — prompt + response group (record)
|
||||
|
||||
``DatasetFactory.load(train_type, load_path, window_size, stride, …)``
|
||||
builds the Store (auto-detecting format) before constructing the
|
||||
matching dataset. Passing ``store=`` skips Store construction.
|
||||
|
||||
When a record dataset (DPO) reads from raw JSONL, a *processor*
|
||||
function (pure ``record -> Dict[str, Tensor]``) is forwarded to
|
||||
:class:`JsonlStore` so tokenisation happens on the fly.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Dict, List, Optional
|
||||
from functools import partial
|
||||
from typing import Callable, Dict, List, Optional
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
@@ -13,198 +37,389 @@ from astrai.dataset.storage import (
|
||||
detect_format,
|
||||
)
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
|
||||
def dpo_tokenize(
|
||||
record: dict,
|
||||
tokenizer,
|
||||
max_len: int = 2048,
|
||||
) -> Optional[dict]:
|
||||
"""Tokenize one DPO record into chosen/rejected + masks.
|
||||
|
||||
Applies the tokenizer's chat template so token sequences match the
|
||||
SFT checkpoint's format. Prompt is rendered with
|
||||
``add_generation_prompt=True``; chosen/rejected are appended as a
|
||||
single assistant turn.
|
||||
|
||||
Accepts:
|
||||
|
||||
- Flat: ``{"prompt": str, "chosen": str, "rejected": str}``
|
||||
- Conv: ``{"prompt": [{role, content}, ...], "chosen": [...], ...}``
|
||||
- Legacy: ``{"input": str, "chosen": str, "rejected": str}``
|
||||
|
||||
No packing, no ``position_ids`` — DPO sequences are independent.
|
||||
"""
|
||||
prompt = record.get("prompt") or record.get("input")
|
||||
chosen = record.get("chosen")
|
||||
rejected = record.get("rejected")
|
||||
if prompt is None or chosen is None or rejected is None:
|
||||
return None
|
||||
|
||||
prompt_messages = _to_messages(prompt)
|
||||
chosen_text = _extract_text(chosen)
|
||||
rejected_text = _extract_text(rejected)
|
||||
if chosen_text is None or rejected_text is None:
|
||||
return None
|
||||
chosen_messages = prompt_messages + [{"role": "assistant", "content": chosen_text}]
|
||||
rejected_messages = prompt_messages + [
|
||||
{"role": "assistant", "content": rejected_text}
|
||||
]
|
||||
|
||||
prompt_ids = tokenizer.apply_chat_template(
|
||||
prompt_messages, tokenize=True, add_generation_prompt=True
|
||||
)
|
||||
ch_ids = tokenizer.apply_chat_template(
|
||||
chosen_messages, tokenize=True, add_generation_prompt=False
|
||||
)
|
||||
re_ids = tokenizer.apply_chat_template(
|
||||
rejected_messages, tokenize=True, add_generation_prompt=False
|
||||
)
|
||||
|
||||
full_ch = ch_ids[:max_len]
|
||||
full_re = re_ids[:max_len]
|
||||
|
||||
prompt_len = min(len(prompt_ids), max_len)
|
||||
ch_mask = [0] * prompt_len + [1] * max(0, len(full_ch) - prompt_len)
|
||||
ch_mask = ch_mask[:max_len]
|
||||
re_mask = [0] * prompt_len + [1] * max(0, len(full_re) - prompt_len)
|
||||
re_mask = re_mask[:max_len]
|
||||
|
||||
return {
|
||||
"chosen": full_ch,
|
||||
"rejected": full_re,
|
||||
"chosen_mask": ch_mask,
|
||||
"rejected_mask": re_mask,
|
||||
}
|
||||
|
||||
|
||||
def _to_messages(value) -> list:
|
||||
"""Accept str or conversation list; return message list."""
|
||||
if isinstance(value, str):
|
||||
return [{"role": "user", "content": value}]
|
||||
if isinstance(value, list):
|
||||
return value
|
||||
return [{"role": "user", "content": str(value)}]
|
||||
|
||||
|
||||
def _extract_text(value) -> Optional[str]:
|
||||
"""Accept str or conversation list; return plain text."""
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
if isinstance(value, list):
|
||||
return "".join(m.get("content", "") for m in value if isinstance(m, dict))
|
||||
return None
|
||||
|
||||
|
||||
def dpo_processor(
|
||||
record: dict,
|
||||
tokenizer,
|
||||
max_len: int = 2048,
|
||||
) -> Dict[str, Tensor]:
|
||||
"""DPO processor: wraps :func:`dpo_tokenize` and returns tensors."""
|
||||
result = dpo_tokenize(record, tokenizer, max_len=max_len)
|
||||
if result is None:
|
||||
raise ValueError(f"Malformed DPO record: {list(record.keys())}")
|
||||
return {
|
||||
"chosen": torch.tensor(result["chosen"], dtype=torch.int32),
|
||||
"rejected": torch.tensor(result["rejected"], dtype=torch.int32),
|
||||
"chosen_mask": torch.tensor(result["chosen_mask"], dtype=torch.bool),
|
||||
"rejected_mask": torch.tensor(result["rejected_mask"], dtype=torch.bool),
|
||||
}
|
||||
|
||||
|
||||
def dpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
|
||||
"""Collate variable-length DPO samples into padded 2-D tensors.
|
||||
|
||||
Input: list of dicts, each with:
|
||||
- chosen: [C_i]
|
||||
- rejected: [R_i]
|
||||
- chosen_mask: [C_i]
|
||||
- rejected_mask: [R_i]
|
||||
|
||||
Output (padded to the max length across chosen/rejected within the batch):
|
||||
- chosen: [B, S_max]
|
||||
- rejected: [B, S_max]
|
||||
- chosen_mask: [B, S_max]
|
||||
- rejected_mask: [B, S_max]
|
||||
"""
|
||||
B = len(batch)
|
||||
S_max = max(b["chosen"].size(0) for b in batch)
|
||||
S_max = max(S_max, max(b["rejected"].size(0) for b in batch))
|
||||
|
||||
chosen = torch.zeros(B, S_max, dtype=torch.long)
|
||||
rejected = torch.zeros(B, S_max, dtype=torch.long)
|
||||
chosen_mask = torch.zeros(B, S_max, dtype=torch.bool)
|
||||
rejected_mask = torch.zeros(B, S_max, dtype=torch.bool)
|
||||
|
||||
for i, b in enumerate(batch):
|
||||
c_len = b["chosen"].size(0)
|
||||
r_len = b["rejected"].size(0)
|
||||
chosen[i, :c_len] = b["chosen"]
|
||||
rejected[i, :r_len] = b["rejected"]
|
||||
chosen_mask[i, :c_len] = b["chosen_mask"]
|
||||
rejected_mask[i, :r_len] = b["rejected_mask"]
|
||||
|
||||
return {
|
||||
"chosen": chosen,
|
||||
"rejected": rejected,
|
||||
"chosen_mask": chosen_mask,
|
||||
"rejected_mask": rejected_mask,
|
||||
}
|
||||
|
||||
|
||||
def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
|
||||
"""Collate variable-length GRPO samples into padded 3-D tensors.
|
||||
|
||||
Input: list of dicts, each with:
|
||||
- prompts: [P_i]
|
||||
- responses: list of G tensors, each [R_ij]
|
||||
- masks: list of G tensors, each [R_ij]
|
||||
- rewards: [G]
|
||||
|
||||
Output:
|
||||
- prompts: [B, P_max]
|
||||
- responses: [B, G, R_max]
|
||||
- masks: [B, G, R_max]
|
||||
- rewards: [B, G]
|
||||
"""
|
||||
B = len(batch)
|
||||
G = len(batch[0]["responses"])
|
||||
P_max = max(b["prompts"].size(0) for b in batch)
|
||||
R_max = max(r.size(0) for b in batch for r in b["responses"])
|
||||
|
||||
prompts = torch.zeros(B, P_max, dtype=torch.long)
|
||||
responses = torch.zeros(B, G, R_max, dtype=torch.long)
|
||||
masks = torch.zeros(B, G, R_max, dtype=torch.bool)
|
||||
rewards = torch.zeros(B, G, dtype=torch.float32)
|
||||
|
||||
for i, b in enumerate(batch):
|
||||
p_len = b["prompts"].size(0)
|
||||
prompts[i, :p_len] = b["prompts"]
|
||||
rewards[i, : b["rewards"].size(0)] = b["rewards"]
|
||||
for g in range(min(G, len(b["responses"]))):
|
||||
r_len = b["responses"][g].size(0)
|
||||
responses[i, g, :r_len] = b["responses"][g]
|
||||
if g < len(b["masks"]):
|
||||
masks[i, g, :r_len] = b["masks"][g]
|
||||
|
||||
return {
|
||||
"prompts": prompts,
|
||||
"responses": responses,
|
||||
"masks": masks,
|
||||
"rewards": rewards,
|
||||
}
|
||||
|
||||
|
||||
def validate_keys(store: Store, required: List[str]) -> None:
|
||||
"""Raise ``KeyError`` if *store* is missing any *required* key."""
|
||||
if not required:
|
||||
return
|
||||
actual = set(store.keys)
|
||||
missing = [k for k in required if k not in actual]
|
||||
if missing:
|
||||
raise KeyError(
|
||||
f"Store at {getattr(store, '_load_path', '?')} is missing required "
|
||||
f"keys {missing}; available keys are {sorted(actual)}."
|
||||
)
|
||||
|
||||
|
||||
class BaseDataset(Dataset, ABC):
|
||||
"""Abstract base class for all dataset types.
|
||||
"""Abstract base class for dataset types.
|
||||
|
||||
Implements common functionality for window-based data fetching.
|
||||
Uses a storage abstraction for format-agnostic data loading.
|
||||
Holds a :class:`Store`. All sample-id indexing is delegated to the
|
||||
store — this class exposes ``__len__`` as ``len(store)`` and the
|
||||
``keys`` property as ``store.keys``. Subclasses implement
|
||||
``__getitem__`` with the train-type-specific key mapping and any
|
||||
training-only index arithmetic (e.g. the next-token ``+1`` shift).
|
||||
"""
|
||||
|
||||
def __init__(self, window_size: int, stride: int):
|
||||
required_keys: List[str] = []
|
||||
|
||||
def __init__(self, store: Store):
|
||||
super().__init__()
|
||||
self.window_size = window_size
|
||||
self.stride = stride
|
||||
self.storage: Optional[Store] = None
|
||||
self.store: Store = store
|
||||
validate_keys(store, self.required_keys)
|
||||
|
||||
@property
|
||||
def required_keys(self) -> List[str]:
|
||||
"""Return required storage keys for this dataset type.
|
||||
|
||||
Subclasses should override to specify expected keys.
|
||||
"""
|
||||
return []
|
||||
|
||||
def _validate_keys(self):
|
||||
if not self.required_keys:
|
||||
return
|
||||
actual_keys = set(self.storage.keys)
|
||||
missing = [k for k in self.required_keys if k not in actual_keys]
|
||||
if missing:
|
||||
raise KeyError(
|
||||
f"Dataset {type(self).__name__} requires keys {self.required_keys}, "
|
||||
f"but storage at {self._load_path} only has {sorted(actual_keys)}. "
|
||||
f"Missing: {missing}"
|
||||
)
|
||||
|
||||
def load(self, load_path: str, storage_type: Optional[str] = None):
|
||||
"""Load dataset from the given path.
|
||||
|
||||
Auto-detects the storage format if not specified.
|
||||
|
||||
Args:
|
||||
load_path: Path to the data directory or file
|
||||
storage_type: Force a specific storage type ("h5", "bin"),
|
||||
or None for auto-detection
|
||||
|
||||
Raises:
|
||||
KeyError: If the loaded storage is missing required keys.
|
||||
"""
|
||||
if storage_type is None:
|
||||
storage_type = detect_format(load_path)
|
||||
self.storage = StoreFactory.create(storage_type)
|
||||
self._load_path = load_path
|
||||
self.storage.load(load_path)
|
||||
self._validate_keys()
|
||||
|
||||
@property
|
||||
def count(self) -> int:
|
||||
"""Return the total number of raw elements (tokens) in the dataset."""
|
||||
if self.storage is None:
|
||||
return 0
|
||||
return len(self.storage)
|
||||
def __len__(self) -> int:
|
||||
return len(self.store)
|
||||
|
||||
@property
|
||||
def keys(self) -> List[str]:
|
||||
"""Return the available data keys."""
|
||||
if self.storage is None:
|
||||
return []
|
||||
return self.storage.keys
|
||||
return self.store.keys
|
||||
|
||||
def get_index(self, index: int) -> tuple:
|
||||
"""Calculate begin and end indices for a sample.
|
||||
|
||||
Args:
|
||||
index: Sample index
|
||||
|
||||
Returns:
|
||||
Tuple of (begin_idx, end_idx)
|
||||
"""
|
||||
if self.storage is None:
|
||||
raise RuntimeError("Dataset not loaded, call load() first")
|
||||
total = len(self.storage)
|
||||
if total <= self.window_size:
|
||||
raise ValueError(
|
||||
f"Data too short: {total} tokens <= window_size {self.window_size}"
|
||||
)
|
||||
|
||||
begin_idx = min(index * self.stride, total - 1 - self.window_size)
|
||||
end_idx = min(begin_idx + self.window_size, total - 1)
|
||||
|
||||
return begin_idx, end_idx
|
||||
@property
|
||||
def token_count(self) -> int:
|
||||
return self.store.token_count
|
||||
|
||||
@abstractmethod
|
||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||
"""Get a single sample by index.
|
||||
|
||||
Must be implemented by subclasses.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def __len__(self) -> int:
|
||||
if self.storage is None:
|
||||
return 0
|
||||
total = len(self.storage)
|
||||
if total <= self.window_size:
|
||||
return 0
|
||||
return (total - 1 - self.window_size) // self.stride + 1
|
||||
|
||||
|
||||
class DatasetFactory(BaseFactory["BaseDataset"]):
|
||||
"""Factory class for creating dataset instances.
|
||||
"""Factory for creating dataset instances by train-type.
|
||||
|
||||
Supports decorator-based registration for extensible dataset types.
|
||||
All default dataset types (seq, sft, dpo, grpo) are registered automatically
|
||||
when their classes are defined with the decorator.
|
||||
|
||||
Example usage:
|
||||
@DatasetFactory.register("custom")
|
||||
class CustomDataset(BaseDataset):
|
||||
...
|
||||
|
||||
dataset = DatasetFactory.create("custom", window_size, stride)
|
||||
Use :meth:`DatasetFactory.register("custom")` to register new
|
||||
dataset classes; they must inherit from :class:`BaseDataset`.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def load(
|
||||
cls,
|
||||
train_type: str,
|
||||
load_path: str,
|
||||
window_size: int,
|
||||
load_path: Optional[str] = None,
|
||||
window_size: int = 0,
|
||||
stride: Optional[int] = None,
|
||||
storage_type: Optional[str] = None,
|
||||
tokenizer_path: Optional[str] = None,
|
||||
max_len: int = 2048,
|
||||
store: Optional[Store] = None,
|
||||
**kwargs,
|
||||
) -> "BaseDataset":
|
||||
"""Create and load a dataset in one step.
|
||||
|
||||
Two entry points:
|
||||
|
||||
- **store given**: bind it directly — the caller fully controls
|
||||
Store construction and processor setup. *load_path*,
|
||||
*storage_type*, *tokenizer_path*, *window_size*, *stride* are
|
||||
ignored.
|
||||
- **store is None**: build a Store from *load_path*, auto-detecting
|
||||
format and constructing a processor when *tokenizer_path* is
|
||||
given for a record dataset on JSONL.
|
||||
|
||||
Args:
|
||||
train_type: Type of training dataset
|
||||
load_path: Path to the data file
|
||||
window_size: Window size for data sampling
|
||||
stride: Stride between consecutive samples (default: same as window_size)
|
||||
storage_type: Storage type ("h5", "bin") or None for auto-detection
|
||||
train_type: Registered dataset name ("seq", "sft", "dpo",
|
||||
"grpo", …).
|
||||
load_path: Path to the data file or directory (ignored if
|
||||
*store* is given).
|
||||
window_size: Stream window length — only meaningful for
|
||||
stream datasets (SEQ/SFT). Record datasets ignore it.
|
||||
stride: Stride between consecutive stream samples
|
||||
(default: same as *window_size*).
|
||||
storage_type: Storage backend ("h5", "bin", "jsonl") or
|
||||
None for auto-detection.
|
||||
tokenizer_path: Path to tokenizer for lazy JSONL
|
||||
tokenisation (record datasets only).
|
||||
max_len: Max sequence length forwarded to processors.
|
||||
store: Pre-built, already-loaded Store instance.
|
||||
**kwargs: Extra arguments forwarded to ``store.load()``.
|
||||
|
||||
Returns:
|
||||
Loaded dataset instance
|
||||
Loaded dataset instance.
|
||||
"""
|
||||
if store is not None:
|
||||
return cls.create(train_type, store=store)
|
||||
|
||||
if load_path is None:
|
||||
raise ValueError("Either load_path or store must be provided")
|
||||
|
||||
if storage_type is None:
|
||||
storage_type = detect_format(load_path)
|
||||
|
||||
if stride is None:
|
||||
stride = window_size
|
||||
|
||||
dataset = cls.create(train_type, window_size, stride)
|
||||
dataset.load(load_path, storage_type=storage_type)
|
||||
processor = cls._maybe_build_processor(
|
||||
train_type, storage_type, tokenizer_path, max_len
|
||||
)
|
||||
|
||||
return dataset
|
||||
store_window = cls._store_window_for(train_type, window_size)
|
||||
store = StoreFactory.create(
|
||||
storage_type,
|
||||
window_size=store_window,
|
||||
stride=stride if stride else store_window,
|
||||
)
|
||||
if processor is not None:
|
||||
store.load(load_path, processor=processor, **kwargs)
|
||||
else:
|
||||
store.load(load_path, **kwargs)
|
||||
|
||||
return cls.create(train_type, store=store)
|
||||
|
||||
@staticmethod
|
||||
def _store_window_for(train_type: str, window_size: int) -> int:
|
||||
"""Stream datasets consume ``window_size``; record datasets ignore it.
|
||||
|
||||
Record datasets (dpo/grpo) treat each record as an independent
|
||||
training unit and never window, so the store is built with
|
||||
``window_size=0`` and ``len(store)`` returns the record count.
|
||||
"""
|
||||
if train_type in ("seq", "sft"):
|
||||
return window_size
|
||||
return 0
|
||||
|
||||
@staticmethod
|
||||
def _maybe_build_processor(
|
||||
train_type: str,
|
||||
storage_type: str,
|
||||
tokenizer_path: Optional[str],
|
||||
max_len: int,
|
||||
) -> Optional[Callable[[dict], Dict[str, Tensor]]]:
|
||||
"""Build an on-the-fly tokenisation processor if applicable.
|
||||
|
||||
Only raw JSONL + record datasets (DPO/GRPO) need a processor;
|
||||
pre-tokenised backends (H5/bin) and stream datasets (SEQ/SFT)
|
||||
return ``None`` so no tokenizer is loaded.
|
||||
"""
|
||||
if tokenizer_path is None or storage_type != "jsonl":
|
||||
return None
|
||||
if train_type == "dpo":
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
|
||||
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
|
||||
return None
|
||||
|
||||
|
||||
@DatasetFactory.register("seq")
|
||||
class SEQDataset(BaseDataset):
|
||||
"""Dataset for sequential next-token prediction training."""
|
||||
"""Dataset for sequential next-token prediction training.
|
||||
|
||||
@property
|
||||
def required_keys(self) -> List[str]:
|
||||
return ["sequence"]
|
||||
Stream mode: ``store.fetch(begin, end, "sequence")`` returns the
|
||||
input window; the +1 shifted call returns the next-token target.
|
||||
"""
|
||||
|
||||
def _fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
||||
return self.storage.fetch(begin_idx, end_idx, "sequence")
|
||||
required_keys = ["sequence"]
|
||||
|
||||
def __getitem__(self, index):
|
||||
begin_idx, end_idx = self.get_index(index)
|
||||
|
||||
x = self._fetch_data(begin_idx, end_idx).to(dtype=torch.long)
|
||||
y = self._fetch_data(begin_idx + 1, end_idx + 1).to(dtype=torch.long)
|
||||
|
||||
return {"input_ids": x, "target_ids": y}
|
||||
def __getitem__(self, index: int):
|
||||
begin, end = self.store.sample_window(index)
|
||||
x = self.store.fetch(begin, end, "sequence")
|
||||
y = self.store.fetch(begin + 1, end + 1, "sequence")
|
||||
return {
|
||||
"input_ids": x.to(dtype=torch.long),
|
||||
"target_ids": y.to(dtype=torch.long),
|
||||
}
|
||||
|
||||
|
||||
@DatasetFactory.register("sft")
|
||||
class SFTDataset(BaseDataset):
|
||||
"""Dataset for supervised fine-tuning with loss masking."""
|
||||
"""Dataset for supervised fine-tuning with loss masking.
|
||||
|
||||
@property
|
||||
def required_keys(self) -> List[str]:
|
||||
return ["sequence", "loss_mask", "position_ids"]
|
||||
Stream mode: ``sequence``/``loss_mask``/``position_ids`` are sliced
|
||||
to the window. ``loss_mask`` and ``target_ids`` use the +1 shifted
|
||||
slice so they align with the predicted positions.
|
||||
"""
|
||||
|
||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||
return self.storage.fetch(begin_idx, end_idx, key)
|
||||
|
||||
def __getitem__(self, index):
|
||||
begin_idx, end_idx = self.get_index(index)
|
||||
|
||||
x = self._fetch_data(begin_idx, end_idx, "sequence")
|
||||
y = self._fetch_data(begin_idx + 1, end_idx + 1, "sequence")
|
||||
position_ids = self._fetch_data(begin_idx, end_idx, "position_ids")
|
||||
loss_mask = self._fetch_data(begin_idx + 1, end_idx + 1, "loss_mask")
|
||||
required_keys = ["sequence", "loss_mask", "position_ids"]
|
||||
|
||||
def __getitem__(self, index: int):
|
||||
begin, end = self.store.sample_window(index)
|
||||
x = self.store.fetch(begin, end, "sequence")
|
||||
y = self.store.fetch(begin + 1, end + 1, "sequence")
|
||||
position_ids = self.store.fetch(begin, end, "position_ids")
|
||||
loss_mask = self.store.fetch(begin + 1, end + 1, "loss_mask")
|
||||
return {
|
||||
"input_ids": x.to(dtype=torch.long),
|
||||
"target_ids": y.to(dtype=torch.long),
|
||||
@@ -215,59 +430,65 @@ class SFTDataset(BaseDataset):
|
||||
|
||||
@DatasetFactory.register("dpo")
|
||||
class DPODataset(BaseDataset):
|
||||
"""Dataset for Direct Preference Optimization training."""
|
||||
"""Record-structured dataset for Direct Preference Optimization.
|
||||
|
||||
@property
|
||||
def required_keys(self) -> List[str]:
|
||||
return ["chosen", "rejected", "chosen_mask", "rejected_mask"]
|
||||
Each sample is one preference pair (chosen + rejected) and is an
|
||||
independent training unit — no windowing, stride, or cross-record
|
||||
concatenation. This keeps each sequence self-contained so attention
|
||||
never leaks across preference pairs.
|
||||
|
||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||
return self.storage.fetch(begin_idx, end_idx, key)
|
||||
Two loading paths (handled by :class:`DatasetFactory`):
|
||||
|
||||
def __getitem__(self, index: int):
|
||||
begin_idx, end_idx = self.get_index(index)
|
||||
- **Pre-tokenized** (H5/bin): ``store.load(path)`` reads per-record
|
||||
tensors; ``__getitem__`` returns them directly.
|
||||
- **Raw JSONL** (``tokenizer_path=...``): builds a lazy processor
|
||||
via :func:`dpo_processor` that tokenises on the fly — no packing,
|
||||
no ``position_ids``.
|
||||
"""
|
||||
|
||||
chosen = self._fetch_data(begin_idx, end_idx, "chosen").to(dtype=torch.long)
|
||||
rejected = self._fetch_data(begin_idx, end_idx, "rejected").to(dtype=torch.long)
|
||||
chosen_mask = self._fetch_data(begin_idx, end_idx, "chosen_mask").to(
|
||||
dtype=torch.bool
|
||||
)
|
||||
rejected_mask = self._fetch_data(begin_idx, end_idx, "rejected_mask").to(
|
||||
dtype=torch.bool
|
||||
)
|
||||
required_keys = ["chosen", "rejected", "chosen_mask", "rejected_mask"]
|
||||
|
||||
def make_processor(self, tokenizer, max_len: int):
|
||||
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
|
||||
|
||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||
return {
|
||||
"chosen": chosen,
|
||||
"rejected": rejected,
|
||||
"chosen_mask": chosen_mask,
|
||||
"rejected_mask": rejected_mask,
|
||||
"chosen": self.store.fetch_record(index, "chosen").to(dtype=torch.long),
|
||||
"rejected": self.store.fetch_record(index, "rejected").to(dtype=torch.long),
|
||||
"chosen_mask": self.store.fetch_record(index, "chosen_mask").to(
|
||||
dtype=torch.bool
|
||||
),
|
||||
"rejected_mask": self.store.fetch_record(index, "rejected_mask").to(
|
||||
dtype=torch.bool
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
@DatasetFactory.register("grpo")
|
||||
class GRPODataset(BaseDataset):
|
||||
"""Dataset for Group Relative Policy Optimization training."""
|
||||
"""Dataset for offline Group Relative Policy Optimization.
|
||||
|
||||
@property
|
||||
def required_keys(self) -> List[str]:
|
||||
return ["prompts", "responses", "masks", "rewards"]
|
||||
Each sample is one prompt with its group of responses and scalar
|
||||
rewards — an independent training unit with no windowing or stride.
|
||||
|
||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||
return self.storage.fetch(begin_idx, end_idx, key)
|
||||
Expected storage layout (produced by JsonlStore or pre-tokenized):
|
||||
|
||||
- ``prompts``: List[Tensor] — one 1-D token tensor per record
|
||||
- ``responses``: List[List[Tensor]] — G response tensors per record
|
||||
- ``masks``: List[List[Tensor]] — G mask tensors per record
|
||||
- ``rewards``: List[Tensor] — one 1-D float tensor (len G) per record
|
||||
"""
|
||||
|
||||
required_keys = ["prompts", "responses", "masks", "rewards"]
|
||||
|
||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||
begin_idx, end_idx = self.get_index(index)
|
||||
|
||||
prompts = self._fetch_data(begin_idx, end_idx, "prompts").to(dtype=torch.long)
|
||||
responses = self._fetch_data(begin_idx, end_idx, "responses").to(
|
||||
dtype=torch.long
|
||||
)
|
||||
masks = self._fetch_data(begin_idx, end_idx, "masks").to(dtype=torch.bool)
|
||||
rewards = self._fetch_data(begin_idx, end_idx, "rewards")
|
||||
|
||||
prompts = self.store.fetch_record(index, "prompts")
|
||||
responses = self.store.fetch_record(index, "responses")
|
||||
masks = self.store.fetch_record(index, "masks")
|
||||
rewards = self.store.fetch_record(index, "rewards")
|
||||
return {
|
||||
"prompts": prompts,
|
||||
"responses": responses,
|
||||
"masks": masks,
|
||||
"rewards": rewards,
|
||||
"prompts": prompts.to(dtype=torch.long),
|
||||
"responses": [r.to(dtype=torch.long) for r in responses],
|
||||
"masks": [m.to(dtype=torch.bool) for m in masks],
|
||||
"rewards": rewards.to(dtype=torch.float32),
|
||||
}
|
||||
|
||||
@@ -5,7 +5,15 @@ import torch.distributed as dist
|
||||
from torch.utils.data import Dataset, Sampler
|
||||
|
||||
|
||||
class ResumableDistributedSampler(Sampler[int]):
|
||||
class RDSampler(Sampler[int]):
|
||||
"""Resumable Distributed Sampler.
|
||||
|
||||
A distributed sampler that supports checkpoint-based resume: iteration
|
||||
state (epoch, position) is tracked so training can continue from the
|
||||
exact sample after a restart. Shards the dataset across
|
||||
``dist.world_size`` replicas with optional shuffling.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
data_source: Dataset,
|
||||
@@ -74,6 +82,7 @@ class ResumableDistributedSampler(Sampler[int]):
|
||||
|
||||
self.epoch += 1
|
||||
self._indices = None
|
||||
self.iter = self.iter % self.num_samples_per_replica
|
||||
|
||||
@property
|
||||
def _remaining(self):
|
||||
|
||||
+520
-150
@@ -1,98 +1,70 @@
|
||||
"""Storage backends for different data formats.
|
||||
|
||||
Layers:
|
||||
- I/O layer: save_* / load_* functions, read/write raw files (HDF5/bin)
|
||||
return Dict[str, List[Tensor]] — format-specific, no state
|
||||
- Store (ABC): central abstraction, normalizes multi-segment into
|
||||
Dict[str, List[Tensor]] per key via _normalize(),
|
||||
fetch() uses bisect across segments — no forced concat
|
||||
- Dataset layer: BaseDataset owns a Store, only calls store.fetch(begin, end, key)
|
||||
Architecture (composition over inheritance):
|
||||
|
||||
Key properties:
|
||||
- Multi-segment: segments kept as-is, no forced concatenation — safe for
|
||||
datasets larger than RAM
|
||||
- Explicit length: _length = min(total elements across keys), set at load,
|
||||
__len__ returns O(1)
|
||||
- Zero-copy mmap: MmapStore wraps np.memmap(mode="r"), all DataLoader
|
||||
workers share OS page-cache pages
|
||||
Store (ABC) — owns _data/_cum/_offsets bookkeeping
|
||||
+ window_size/stride for sample-id
|
||||
indexing. __getitem__/__len__ produce
|
||||
the smallest iterable unit so Dataset
|
||||
classes are pure delegators.
|
||||
Streamable (mixin) — raw token slice fetch(begin, end, keys)
|
||||
Recordable (mixin) — raw record slice fetch_record(idx, keys)
|
||||
|
||||
H5Store(Store, Streamable, Recordable)
|
||||
MmapStore(Store, Streamable, Recordable)
|
||||
JsonlStore(Store, Streamable, Recordable)
|
||||
|
||||
Each mixin is a stateless trait that relies on ``self._data`` etc.
|
||||
provided by :class:`Store`. Concrete stores mix in whichever access
|
||||
primitives they support — ``Store`` is the sole base class, so there is
|
||||
no diamond inheritance or MRO ambiguity.
|
||||
|
||||
Sample-id indexing lives on :class:`Store`, not on the dataset:
|
||||
|
||||
- **Stream mode** (``window_size > 0``): ``len(store)`` returns the number
|
||||
of ``(window_size, stride)`` windows that fit in the token river;
|
||||
``store[i]`` returns the *i*-th window as a dict of per-key tensors;
|
||||
``store.sample_window(i)`` exposes the underlying ``(begin, end)``
|
||||
token slice for callers (e.g. next-token trainers) that need a +1
|
||||
shifted companion window.
|
||||
- **Record mode** (``num_records > 0``): ``len(store)`` returns the
|
||||
record count; ``store[i]`` returns the *i*-th record dict.
|
||||
|
||||
Raw token/record access via :meth:`fetch` / :meth:`fetch_record`
|
||||
remains available for low-level callers that want explicit index
|
||||
control. ``store.token_count`` is the total stream token count (what
|
||||
``len(store)`` used to mean in the legacy stream-only API).
|
||||
|
||||
``segments_are_records`` (class attribute on each Store subclass)
|
||||
tells ``_normalize`` whether segments are inherently per-record (H5/
|
||||
JSONL) or opaque shards (bin). Record access for bin relies on
|
||||
``_offsets`` instead.
|
||||
|
||||
:class:`JsonlStore` supports a lazy mode (``processor=fn``) that keeps
|
||||
raw records and defers tokenisation to ``fetch_record`` — used by DPO
|
||||
to train directly from a ``.jsonl`` file without a pre-tokenised copy.
|
||||
"""
|
||||
|
||||
import bisect
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Union
|
||||
from typing import Callable, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import h5py
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.preprocessing.transform import TokenizeTransform
|
||||
from astrai.serialization import (
|
||||
load_bin,
|
||||
load_bin_offsets,
|
||||
load_h5,
|
||||
)
|
||||
|
||||
|
||||
def save_h5(file_path: str, file_name: str, tensor_group: Dict[str, List[Tensor]]):
|
||||
os.makedirs(file_path, exist_ok=True)
|
||||
full_file_path = os.path.join(file_path, f"{file_name}.h5")
|
||||
with h5py.File(full_file_path, "w") as f:
|
||||
for key, tensors in tensor_group.items():
|
||||
grp = f.create_group(key)
|
||||
for idx, tensor in enumerate(tensors):
|
||||
arr = tensor.cpu().numpy()
|
||||
grp.create_dataset(f"data_{idx}", data=arr)
|
||||
|
||||
|
||||
def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
|
||||
tensor_group: Dict[str, List[Tensor]] = {}
|
||||
|
||||
root_path = Path(file_path)
|
||||
h5_files = list(root_path.rglob("*.h5")) + list(root_path.rglob("*.hdf5"))
|
||||
|
||||
for h5_file in h5_files:
|
||||
with h5py.File(h5_file, "r") as f:
|
||||
for key in f.keys():
|
||||
grp = f[key]
|
||||
dsets = []
|
||||
for dset_name in grp.keys():
|
||||
dset = grp[dset_name]
|
||||
tensor = torch.from_numpy(dset[:])
|
||||
if share_memory:
|
||||
tensor = tensor.share_memory_()
|
||||
dsets.append(tensor)
|
||||
|
||||
if tensor_group.get(key) is None:
|
||||
tensor_group[key] = []
|
||||
tensor_group[key].extend(dsets)
|
||||
|
||||
return tensor_group
|
||||
|
||||
|
||||
def save_bin(file_path: str, tensor_group: Dict[str, List[Tensor]]):
|
||||
os.makedirs(file_path, exist_ok=True)
|
||||
meta = {}
|
||||
for key, tensors in tensor_group.items():
|
||||
cat = torch.cat(tensors, dim=0)
|
||||
meta[key] = {"shape": list(cat.shape), "dtype": str(cat.dtype).split(".")[-1]}
|
||||
np.asarray(cat.cpu().numpy()).tofile(os.path.join(file_path, f"{key}.bin"))
|
||||
with open(os.path.join(file_path, "meta.json"), "w") as f:
|
||||
json.dump(meta, f)
|
||||
|
||||
|
||||
def load_bin(file_path: str) -> Dict[str, List[Tensor]]:
|
||||
with open(os.path.join(file_path, "meta.json"), "r") as f:
|
||||
meta = json.load(f)
|
||||
segments: Dict[str, List[Tensor]] = {}
|
||||
for key, info in meta.items():
|
||||
arr = np.memmap(
|
||||
os.path.join(file_path, f"{key}.bin"),
|
||||
dtype=info["dtype"],
|
||||
mode="r+",
|
||||
shape=tuple(info["shape"]),
|
||||
)
|
||||
segments[key] = [torch.from_numpy(arr)]
|
||||
return segments
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def detect_format(load_path: str) -> str:
|
||||
@@ -102,7 +74,7 @@ def detect_format(load_path: str) -> str:
|
||||
load_path: Directory or file path
|
||||
|
||||
Returns:
|
||||
Format string ("h5" or "bin")
|
||||
Format string ("h5", "bin", "jsonl", or "processed")
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If no supported data files are found
|
||||
@@ -112,6 +84,8 @@ def detect_format(load_path: str) -> str:
|
||||
suffix = root.suffix.lower()
|
||||
if suffix in (".h5", ".hdf5"):
|
||||
return "h5"
|
||||
if suffix == ".jsonl":
|
||||
return "jsonl"
|
||||
raise ValueError(f"Unsupported file format: {suffix}")
|
||||
|
||||
h5_files = [
|
||||
@@ -128,139 +102,535 @@ def detect_format(load_path: str) -> str:
|
||||
) > 0
|
||||
if has_meta:
|
||||
return "bin"
|
||||
jsonl_files = [
|
||||
Path(p) for p in glob.glob(str(root / "**" / "*.jsonl"), recursive=True)
|
||||
]
|
||||
if jsonl_files:
|
||||
return "jsonl"
|
||||
raise FileNotFoundError(f"No supported data files found at {load_path}")
|
||||
|
||||
|
||||
class Store(ABC):
|
||||
"""String keys -> segmented tensors with ``fetch(begin, end, keys)``.
|
||||
"""Common base for all storage backends.
|
||||
|
||||
Each key maps to one or more tensor segments (no forced concatenation).
|
||||
``len(store)`` returns ``self._length`` (explicit, O(1)), the minimum
|
||||
total element count across all keys.
|
||||
A Store owns both its data layout AND its sample-id → token/record
|
||||
index translation. Datasets are thin wrappers that bind a Store
|
||||
to a particular train-type's key mapping; they never know about
|
||||
window/stride math.
|
||||
|
||||
Subclasses fill ``self._data`` and ``self._cum`` during ``load()``
|
||||
via ``_normalize()``.
|
||||
Two iteration modes:
|
||||
|
||||
- **Stream** (``window_size > 0``): data is treated as one long
|
||||
token river. ``len(store)`` returns the number of windows;
|
||||
``store[i]`` slices every stream-compatible key to window ``i``;
|
||||
``store.sample_window(i)`` returns the ``(begin, end)`` token
|
||||
slice for callers needing a +1 shifted companion window.
|
||||
- **Record** (``num_records > 0``): data is per-record.
|
||||
``len(store)`` returns ``num_records``; ``store[i]`` returns
|
||||
the *i*-th record as a dict.
|
||||
|
||||
Raw token slicing is still available via :meth:`fetch` (mixed in
|
||||
by :class:`Streamable`) when a store has stream support configured.
|
||||
Raw record slicing via :meth:`fetch_record` (mixed in by
|
||||
:class:`Recordable`) when a store has record support.
|
||||
|
||||
``token_count`` exposes the raw total stream length — this is what
|
||||
``len(store)`` returned in the legacy stream-only API and what
|
||||
stream-bound ``fetch`` uses for its bounds check.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
segments_are_records: bool = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
window_size: int = 0,
|
||||
stride: Optional[int] = None,
|
||||
):
|
||||
self._data: Dict[str, List[Tensor]] = {}
|
||||
self._cum: Dict[str, List[int]] = {}
|
||||
self._offsets: Dict[str, List[int]] = {}
|
||||
self._length: int = 0
|
||||
self._num_records: int = 0
|
||||
self._window_size: int = int(window_size)
|
||||
self._stride: int = int(stride) if stride is not None else int(window_size)
|
||||
|
||||
@abstractmethod
|
||||
def load(self, path: str) -> None:
|
||||
def load(self, path: str, **kwargs) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
@property
|
||||
def keys(self) -> List[str]:
|
||||
return list(self._data.keys())
|
||||
|
||||
def __len__(self) -> int:
|
||||
@property
|
||||
def window_size(self) -> int:
|
||||
return self._window_size
|
||||
|
||||
@property
|
||||
def stride(self) -> int:
|
||||
return self._stride
|
||||
|
||||
@property
|
||||
def token_count(self) -> int:
|
||||
"""Total tokens across all stream segments.
|
||||
|
||||
Useful for the bounds-checked raw :meth:`fetch` and as the
|
||||
legacy ``len(store)`` value.
|
||||
"""
|
||||
return self._length
|
||||
|
||||
@property
|
||||
def num_records(self) -> int:
|
||||
"""Number of records available via :meth:`fetch_record`.
|
||||
|
||||
Non-zero only when the backing layout provides per-record
|
||||
indexing (H5/JSONL segments or bin ``_offsets``).
|
||||
"""
|
||||
return self._num_records
|
||||
|
||||
@property
|
||||
def num_samples(self) -> int:
|
||||
"""Number of items produced by ``__getitem__``.
|
||||
|
||||
Stream-mode wins when ``window_size > 0`` and there are tokens
|
||||
to slice; otherwise falls back to ``num_records``.
|
||||
"""
|
||||
if self._window_size > 0 and self._length > 0:
|
||||
total = self._length
|
||||
w = self._window_size
|
||||
if total <= w:
|
||||
return 0
|
||||
return (total - 1 - w) // self._stride + 1
|
||||
return self._num_records
|
||||
|
||||
def __len__(self) -> int:
|
||||
return self.num_samples
|
||||
|
||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||
if index < 0:
|
||||
index += self.num_samples
|
||||
if not 0 <= index < self.num_samples:
|
||||
raise IndexError(
|
||||
f"Store index out of range: {index}, num_samples={self.num_samples}"
|
||||
)
|
||||
if self._window_size > 0 and self._length > 0:
|
||||
begin, end = self.sample_window(index)
|
||||
keys = self._stream_keys()
|
||||
return {k: self.fetch(begin, end, k) for k in keys}
|
||||
return self.fetch_record(index, self._record_keys())
|
||||
|
||||
def sample_window(self, index: int) -> Tuple[int, int]:
|
||||
"""Return ``(begin, end)`` token positions for stream sample *index*.
|
||||
|
||||
The clipped tail keeps the last reachable window inside the
|
||||
token river instead of overshooting. Caller is responsible
|
||||
for staying within :attr:`num_samples`: an out-of-range index
|
||||
raises ``IndexError``.
|
||||
"""
|
||||
if self._window_size <= 0:
|
||||
raise RuntimeError("sample_window() requires window_size > 0 (stream mode)")
|
||||
if self._window_size <= 0 or self._length <= self._window_size:
|
||||
raise IndexError(
|
||||
f"Data too short for window: token_count={self._length}, "
|
||||
f"window_size={self._window_size}"
|
||||
)
|
||||
if not 0 <= index < self.num_samples:
|
||||
raise IndexError(
|
||||
f"Sample index out of range: {index}, num_samples={self.num_samples}"
|
||||
)
|
||||
total = self._length
|
||||
begin = min(index * self._stride, total - 1 - self._window_size)
|
||||
end = min(begin + self._window_size, total - 1)
|
||||
return begin, end
|
||||
|
||||
def _stream_keys(self) -> List[str]:
|
||||
out: List[str] = []
|
||||
for k, tensors in self._data.items():
|
||||
if tensors and isinstance(tensors[0], list):
|
||||
continue
|
||||
out.append(k)
|
||||
return out
|
||||
|
||||
def _record_keys(self) -> List[str]:
|
||||
return list(self._data.keys())
|
||||
|
||||
def _normalize(
|
||||
self,
|
||||
raw: Dict[str, list],
|
||||
offsets: Optional[Dict[str, List[int]]] = None,
|
||||
):
|
||||
"""Register segments and pre-compute indices for both access modes.
|
||||
|
||||
Stream mode: ``_cum[key]`` accumulates per-segment lengths so
|
||||
``Streamable._fetch_stream_key`` can bisect across segments
|
||||
without concatenation.
|
||||
|
||||
Record mode: if *offsets* is provided (bin layout),
|
||||
``_offsets[key]`` stores cumulative per-record offsets into the
|
||||
single concatenated segment. Otherwise, when
|
||||
``segments_are_records`` is True (H5/JSONL), ``_data[key]`` is
|
||||
a per-record list and ``fetch_record`` indexes it directly.
|
||||
|
||||
Nested keys (GRPO ``responses``/``masks`` as
|
||||
``List[List[Tensor]]``) are stored as-is and excluded from both
|
||||
cumulative bookkeepings — they are only accessed record-by-record.
|
||||
"""
|
||||
flat_lengths = []
|
||||
for key, tensors in raw.items():
|
||||
self._data[key] = tensors
|
||||
if not tensors:
|
||||
self._cum[key] = []
|
||||
flat_lengths.append(0)
|
||||
continue
|
||||
if isinstance(tensors[0], list):
|
||||
self._cum[key] = []
|
||||
continue
|
||||
cum = []
|
||||
total = 0
|
||||
for t in tensors:
|
||||
total += t.shape[0]
|
||||
cum.append(total)
|
||||
self._cum[key] = cum
|
||||
flat_lengths.append(cum[-1] if cum else 0)
|
||||
self._length = min(flat_lengths) if flat_lengths else 0
|
||||
|
||||
valid_offsets: Dict[str, List[int]] = {}
|
||||
if offsets:
|
||||
for key, off in offsets.items():
|
||||
segs = self._data.get(key, [])
|
||||
if len(segs) == 1 and len(off) > 1:
|
||||
valid_offsets[key] = off
|
||||
elif len(segs) > 1:
|
||||
logger.warning(
|
||||
"Key '%s' has %d segments with offsets — record mode "
|
||||
"disabled for this key (multi-shard bin+offsets not "
|
||||
"supported). Merge shards or use H5/JSONL.",
|
||||
key,
|
||||
len(segs),
|
||||
)
|
||||
self._offsets = valid_offsets
|
||||
if valid_offsets:
|
||||
record_counts = [len(v) - 1 for v in valid_offsets.values()]
|
||||
self._num_records = min(record_counts) if record_counts else 0
|
||||
elif self.segments_are_records:
|
||||
per_record_counts = []
|
||||
for key, tensors in self._data.items():
|
||||
if tensors and isinstance(tensors[0], list):
|
||||
continue
|
||||
per_record_counts.append(len(tensors))
|
||||
self._num_records = min(per_record_counts) if per_record_counts else 0
|
||||
else:
|
||||
self._num_records = 0
|
||||
|
||||
|
||||
class Streamable:
|
||||
"""Mixin granting raw token-stream access via :meth:`fetch`.
|
||||
|
||||
Stateless trait relying on ``self._data``, ``self._cum``,
|
||||
``self._length`` maintained by :class:`Store`. Stream mode is
|
||||
active when the owning store has ``window_size > 0``; for stores
|
||||
that can also serve record access (H5/JSONL/bin+offsets), the
|
||||
``fetch_record`` API from :class:`Recordable` is used instead.
|
||||
"""
|
||||
|
||||
def fetch(
|
||||
self,
|
||||
begin: int,
|
||||
end: int,
|
||||
keys: Union[str, List[str]],
|
||||
):
|
||||
if not self._data:
|
||||
raise RuntimeError("Store not loaded")
|
||||
if not (0 <= begin < self._length and 0 <= end <= self._length):
|
||||
raise ValueError(
|
||||
f"Index out of bounds: begin={begin}, end={end}, length={self._length}"
|
||||
)
|
||||
if isinstance(keys, str):
|
||||
return self._fetch_key(keys, begin, end)
|
||||
return {k: self._fetch_key(k, begin, end) for k in keys}
|
||||
return _stream_fetch(self, begin, end, keys)
|
||||
|
||||
def _fetch_key(self, key: str, begin: int, end: int) -> Tensor:
|
||||
"""Fetch slice [begin, end) across potentially multiple segments."""
|
||||
segments = self._data[key]
|
||||
cum = self._cum[key]
|
||||
seg_start = bisect.bisect_right(cum, begin)
|
||||
seg_end = bisect.bisect_left(cum, end)
|
||||
|
||||
results = []
|
||||
for i in range(seg_start, seg_end + 1):
|
||||
prev = cum[i - 1] if i > 0 else 0
|
||||
s = max(begin - prev, 0)
|
||||
e = min(end - prev, segments[i].shape[0])
|
||||
results.append(segments[i][s:e])
|
||||
|
||||
return results[0] if len(results) == 1 else torch.cat(results, dim=0)
|
||||
|
||||
def _normalize(self, raw: Dict[str, List[Tensor]]):
|
||||
"""Register segments and pre-compute cumulative lengths.
|
||||
|
||||
Does NOT concatenate — segments are kept as-is to avoid OOM on
|
||||
large datasets. Sets ``self._length`` to the minimum total
|
||||
element count across all keys.
|
||||
"""
|
||||
for key, tensors in raw.items():
|
||||
self._data[key] = tensors
|
||||
cum = []
|
||||
total = 0
|
||||
for t in tensors:
|
||||
total += t.shape[0]
|
||||
cum.append(total)
|
||||
self._cum[key] = cum
|
||||
self._length = (
|
||||
min((cum[-1] if cum else 0) for cum in self._cum.values())
|
||||
if self._cum
|
||||
else 0
|
||||
def _stream_fetch(self, begin: int, end: int, keys: Union[str, List[str]]):
|
||||
if not getattr(self, "_data", None):
|
||||
raise RuntimeError("Store not loaded")
|
||||
if not (0 <= begin < self._length and 0 <= end <= self._length):
|
||||
raise ValueError(
|
||||
f"Index out of bounds: begin={begin}, end={end}, length={self._length}"
|
||||
)
|
||||
if isinstance(keys, str):
|
||||
return _fetch_stream_key(self, keys, begin, end)
|
||||
return {k: _fetch_stream_key(self, k, begin, end) for k in keys}
|
||||
|
||||
|
||||
def _fetch_stream_key(self, key: str, begin: int, end: int) -> Tensor:
|
||||
segments = self._data[key]
|
||||
cum = self._cum[key]
|
||||
seg_start = bisect.bisect_right(cum, begin)
|
||||
seg_end = bisect.bisect_left(cum, end)
|
||||
|
||||
results = []
|
||||
for i in range(seg_start, seg_end + 1):
|
||||
prev = cum[i - 1] if i > 0 else 0
|
||||
s = max(begin - prev, 0)
|
||||
e = min(end - prev, segments[i].shape[0])
|
||||
results.append(segments[i][s:e])
|
||||
|
||||
return results[0] if len(results) == 1 else torch.cat(results, dim=0)
|
||||
|
||||
|
||||
class Recordable:
|
||||
"""Mixin granting raw record access via :meth:`fetch_record`.
|
||||
|
||||
Stateless trait relying on ``self._data``, ``self._offsets``,
|
||||
``self._num_records`` maintained by :class:`Store`.
|
||||
"""
|
||||
|
||||
def fetch_record(
|
||||
self,
|
||||
index: int,
|
||||
keys: Union[str, List[str]],
|
||||
):
|
||||
return _record_fetch(self, index, keys)
|
||||
|
||||
|
||||
def _record_fetch(self, index: int, keys: Union[str, List[str]]):
|
||||
if not getattr(self, "_data", None) and self._num_records == 0:
|
||||
raise RuntimeError("Store not loaded")
|
||||
if not 0 <= index < self._num_records:
|
||||
raise ValueError(
|
||||
f"Record index out of bounds: {index}, num_records={self._num_records}"
|
||||
)
|
||||
if isinstance(keys, str):
|
||||
return _fetch_record_key(self, keys, index)
|
||||
return {k: _fetch_record_key(self, k, index) for k in keys}
|
||||
|
||||
|
||||
def _fetch_record_key(self, key: str, index: int):
|
||||
offsets = self._offsets.get(key)
|
||||
if offsets:
|
||||
start = offsets[index]
|
||||
end = (
|
||||
offsets[index + 1]
|
||||
if index + 1 < len(offsets)
|
||||
else self._data[key][0].shape[0]
|
||||
)
|
||||
return self._data[key][0][start:end]
|
||||
return self._data[key][index]
|
||||
|
||||
|
||||
class StoreFactory(BaseFactory["Store"]):
|
||||
"""Factory for creating Store instances by type name.
|
||||
|
||||
Example::
|
||||
|
||||
@StoreFactory.register("custom")
|
||||
class CustomStore(Store):
|
||||
...
|
||||
"""
|
||||
"""Factory for creating Store instances by type name."""
|
||||
|
||||
|
||||
@StoreFactory.register("h5")
|
||||
class H5Store(Store):
|
||||
"""HDF5-based storage backend (pre-tokenized data)."""
|
||||
class H5Store(Store, Streamable, Recordable):
|
||||
"""HDF5-based storage backend (pre-tokenized data).
|
||||
|
||||
def load(self, path: str):
|
||||
Each key is stored as a group of per-record datasets (``data_0``,
|
||||
``data_1``, …). Supports both access modes:
|
||||
|
||||
- **Stream**: ``fetch(begin, end, key)`` and ``store[i]`` slice
|
||||
across concatenated records via ``_cum`` — used by SEQ/SFT.
|
||||
- **Record**: ``fetch_record(i, key)`` and ``store[i]`` (when
|
||||
``window_size == 0``) index ``_data[key]`` directly — used by
|
||||
DPO/GRPO.
|
||||
"""
|
||||
|
||||
segments_are_records = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
window_size: int = 0,
|
||||
stride: Optional[int] = None,
|
||||
):
|
||||
super().__init__(window_size=window_size, stride=stride)
|
||||
|
||||
def load(self, path: str, **kwargs):
|
||||
self._normalize(load_h5(path))
|
||||
|
||||
|
||||
@StoreFactory.register("bin")
|
||||
class MmapStore(Store):
|
||||
class MmapStore(Store, Streamable, Recordable):
|
||||
"""Memory-mapped binary storage backend.
|
||||
|
||||
Each key is a single .bin file backed by ``np.memmap(mode="r")``.
|
||||
No per-process memory duplication — all DataLoader workers share the
|
||||
same OS page-cache pages.
|
||||
|
||||
Format on disk::
|
||||
Supports both access modes:
|
||||
|
||||
data_root/
|
||||
meta.json # {key: {shape, dtype}, ...}
|
||||
<key>.bin # raw numpy array, one per key
|
||||
- **Stream**: always available via :meth:`fetch`.
|
||||
- **Record** (``fetch_record(i, key)``): only when ``meta.json``
|
||||
contains per-record ``offsets`` (written via
|
||||
``save_bin(..., record_keys=...)``). Legacy bin files without
|
||||
offsets have ``num_records == 0`` and ``len(store)`` reflects the
|
||||
windowed sample count when ``window_size > 0``.
|
||||
|
||||
``segments_are_records`` is ``False`` here (bin segments are
|
||||
contiguous streams, not per-record) — record access is driven
|
||||
purely by ``_offsets``.
|
||||
"""
|
||||
|
||||
def load(self, path: str):
|
||||
segments_are_records = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
window_size: int = 0,
|
||||
stride: Optional[int] = None,
|
||||
):
|
||||
super().__init__(window_size=window_size, stride=stride)
|
||||
self._mmap_refs: List[Tensor] = []
|
||||
|
||||
def load(self, path: str, **kwargs):
|
||||
self._mmap_refs = []
|
||||
root = Path(path)
|
||||
all_raw: Dict[str, List[Tensor]] = {}
|
||||
all_offsets: Dict[str, List[int]] = {}
|
||||
meta_paths = [
|
||||
Path(p) for p in glob.glob(str(root / "**" / "meta.json"), recursive=True)
|
||||
]
|
||||
for meta_path in meta_paths:
|
||||
raw = load_bin(str(meta_path.parent))
|
||||
off = load_bin_offsets(str(meta_path.parent))
|
||||
for key, tensors in raw.items():
|
||||
if key not in all_raw:
|
||||
all_raw[key] = []
|
||||
all_raw[key].extend(tensors)
|
||||
for key, o in off.items():
|
||||
if key not in all_offsets:
|
||||
all_offsets[key] = []
|
||||
all_offsets[key].extend(o)
|
||||
if not meta_paths:
|
||||
raise FileNotFoundError(f"No meta.json found under {path}")
|
||||
self._normalize(all_raw)
|
||||
self._normalize(all_raw, offsets=all_offsets or None)
|
||||
for tensors in self._data.values():
|
||||
self._mmap_refs.extend(tensors)
|
||||
|
||||
|
||||
class JsonlSource:
|
||||
"""Read raw JSON records from a ``.jsonl`` file or directory.
|
||||
|
||||
A thin reader used by :class:`JsonlStore` in processor mode — holds
|
||||
no tokenizer, performs no tokenisation, just yields dicts.
|
||||
"""
|
||||
|
||||
def __init__(self, path: str):
|
||||
self.path = Path(path)
|
||||
self._records: Optional[List[dict]] = None
|
||||
|
||||
def load(self) -> List[dict]:
|
||||
if self._records is None:
|
||||
self._records = self._read(self.path)
|
||||
return self._records
|
||||
|
||||
@staticmethod
|
||||
def _read(root: Path) -> List[dict]:
|
||||
if root.is_file():
|
||||
return JsonlSource._read_file(root)
|
||||
return JsonlSource._read_dir(root)
|
||||
|
||||
@staticmethod
|
||||
def _read_file(path: Path) -> List[dict]:
|
||||
records: List[dict] = []
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
try:
|
||||
records.append(json.loads(line))
|
||||
except json.JSONDecodeError:
|
||||
logger.warning("Failed to parse JSON line in %s, skipping", path)
|
||||
return records
|
||||
|
||||
@staticmethod
|
||||
def _read_dir(root: Path) -> List[dict]:
|
||||
records: List[dict] = []
|
||||
for jsonl_path in sorted(root.glob("*.jsonl")):
|
||||
records.extend(JsonlSource._read_file(jsonl_path))
|
||||
return records
|
||||
|
||||
|
||||
@StoreFactory.register("jsonl")
|
||||
class JsonlStore(Store, Streamable, Recordable):
|
||||
"""JSONL reader with two tokenisation modes.
|
||||
|
||||
A JSONL dataset is a ``.jsonl`` file or a directory of ``*.jsonl``
|
||||
files plus (optionally) a ``dataset_config.json`` describing the
|
||||
tokenization pipeline.
|
||||
|
||||
Two modes, selected at :meth:`load` time:
|
||||
|
||||
- **Eager** (default): applies a :class:`TokenizeTransform` to every
|
||||
record at load time and registers per-key tensors via
|
||||
``_normalize``. Both ``fetch`` (stream) and ``fetch_record``
|
||||
(record) work.
|
||||
- **Lazy** (``processor=fn`` passed): keeps raw records and defers
|
||||
tokenisation to ``fetch_record``. Only record access works —
|
||||
``len(store)`` returns ``num_records``; stream primitives raise.
|
||||
"""
|
||||
|
||||
CONFIG_NAME = "dataset_config.json"
|
||||
segments_are_records = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
window_size: int = 0,
|
||||
stride: Optional[int] = None,
|
||||
):
|
||||
super().__init__(window_size=window_size, stride=stride)
|
||||
self._source: Optional[JsonlSource] = None
|
||||
self._processor: Optional[Callable[[dict], Dict[str, Tensor]]] = None
|
||||
self._keys_cache: Optional[List[str]] = None
|
||||
|
||||
def load(self, path: str, transform=None, processor=None, **kwargs):
|
||||
self._source = JsonlSource(path)
|
||||
records = self._source.load()
|
||||
|
||||
if processor is not None:
|
||||
self._processor = processor
|
||||
self._num_records = len(records)
|
||||
return
|
||||
|
||||
if transform is None:
|
||||
root = Path(path)
|
||||
config_path = root / self.CONFIG_NAME if root.is_dir() else None
|
||||
if config_path is None or not config_path.exists():
|
||||
raise FileNotFoundError(
|
||||
f"JSONL dataset config not found. Expected "
|
||||
f"{self.CONFIG_NAME} alongside *.jsonl files, pass an "
|
||||
f"explicit transform, or pass processor= for lazy "
|
||||
f"on-the-fly tokenisation."
|
||||
)
|
||||
transform = TokenizeTransform.from_config_file(str(config_path))
|
||||
|
||||
transformed = transform.apply(records)
|
||||
self._normalize(transformed)
|
||||
|
||||
@property
|
||||
def keys(self) -> List[str]:
|
||||
if self._processor is not None:
|
||||
if self._keys_cache is None and self._num_records > 0:
|
||||
sample = self._processor(self._source.load()[0])
|
||||
self._keys_cache = list(sample.keys())
|
||||
return self._keys_cache or []
|
||||
return list(self._data.keys())
|
||||
|
||||
def fetch_record(self, index: int, keys: Union[str, List[str]]):
|
||||
if self._processor is not None:
|
||||
if not 0 <= index < self._num_records:
|
||||
raise ValueError(
|
||||
f"Record index out of bounds: {index}, "
|
||||
f"num_records={self._num_records}"
|
||||
)
|
||||
record = self._source.load()[index]
|
||||
data = self._processor(record)
|
||||
if isinstance(keys, str):
|
||||
return data[keys]
|
||||
return {k: data[k] for k in keys}
|
||||
return _record_fetch(self, index, keys)
|
||||
|
||||
def fetch(self, begin: int, end: int, keys: Union[str, List[str]]):
|
||||
if self._processor is not None:
|
||||
raise RuntimeError(
|
||||
"JsonlStore in lazy (processor) mode does not support "
|
||||
"stream fetch(); use fetch_record() instead."
|
||||
)
|
||||
return _stream_fetch(self, begin, end, keys)
|
||||
|
||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||
if self._processor is not None:
|
||||
return self.fetch_record(index, self._record_keys())
|
||||
return super().__getitem__(index)
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
"""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)
|
||||
|
||||
Interface (shared by all wrappers):
|
||||
causal_offset: -1 = non-causal; >=0 = absolute position of first Q token
|
||||
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True = keep)
|
||||
scale: 0.0 = auto (1/sqrt(head_dim)); >0 = explicit
|
||||
layout: "bhld" (default) or "blhd"
|
||||
|
||||
Causal and mask can coexist — both are applied simultaneously.
|
||||
|
||||
Each wrapper dispatches to its compiled CUDA kernel (``astrai.extension.attn_*``)
|
||||
when available, otherwise falls back to ``torch.nn.functional.scaled_dot_product_attention``.
|
||||
"""
|
||||
|
||||
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,246 @@
|
||||
"""GQA attention wrapper functions — one entry point per compiled kernel.
|
||||
|
||||
Each wrapper dispatches to its CUDA kernel (loaded in ``loader.py``) when
|
||||
available, otherwise falls back to ``torch`` SDPA.
|
||||
|
||||
Interface (all functions):
|
||||
causal_offset: -1 = non-causal; >=0 = absolute position of first Q token
|
||||
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool)
|
||||
scale: 0.0 = auto (1/sqrt(head_dim)); >0 = explicit
|
||||
layout: "bhld" (default) or "blhd"
|
||||
|
||||
Add new kernel wrappers here; split into per-variant files only if this file
|
||||
grows large.
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from astrai.extension.loader import _available, _modules
|
||||
|
||||
_LAYOUT_CODES: dict[str, int] = {"bhld": 0, "blhd": 1}
|
||||
|
||||
|
||||
def _parse_layout(layout: str | int) -> int:
|
||||
if isinstance(layout, int):
|
||||
return layout
|
||||
code = _LAYOUT_CODES.get(layout.lower())
|
||||
if code is None:
|
||||
raise ValueError(
|
||||
f"unknown layout '{layout}', expected one of {list(_LAYOUT_CODES)}"
|
||||
)
|
||||
return code
|
||||
|
||||
|
||||
def _to_bhld(t: torch.Tensor, layout: int) -> torch.Tensor:
|
||||
"""Normalize to b h l d view. Zero-copy transpose if layout==1 (b l h d)."""
|
||||
if layout == 1:
|
||||
return t.transpose(1, 2)
|
||||
return t
|
||||
|
||||
|
||||
def _expand_kv_heads(
|
||||
k: torch.Tensor, v: torch.Tensor, q_head: int
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Expand K/V heads to match Q heads for GQA fallback."""
|
||||
kv_head = k.size(1)
|
||||
if kv_head == q_head:
|
||||
return k, v
|
||||
group = q_head // kv_head
|
||||
k = k.repeat_interleave(group, dim=1)
|
||||
v = v.repeat_interleave(group, dim=1)
|
||||
return k, v
|
||||
|
||||
|
||||
def _build_attn_mask(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
mask: torch.Tensor | None,
|
||||
causal_offset: int,
|
||||
scale: float,
|
||||
) -> tuple[torch.Tensor | None, float]:
|
||||
"""Build SDPA-compatible attn_mask + resolved scale.
|
||||
|
||||
q and k must already be in b h l d layout.
|
||||
Causal and mask can coexist: causal sets -inf above the diagonal, mask
|
||||
sets -inf for padded positions. Both are OR'd into a single bool mask.
|
||||
"""
|
||||
q_len = q.size(2)
|
||||
kv_len = k.size(2)
|
||||
head_dim = q.size(3)
|
||||
resolved_scale = scale if scale and scale > 0 else 1.0 / math.sqrt(head_dim)
|
||||
|
||||
attn_mask = None
|
||||
|
||||
if mask is not None:
|
||||
if mask.dim() == 2:
|
||||
# [batch, kv_len] → [batch, 1, 1, kv_len]
|
||||
attn_mask = mask[:, None, None, :]
|
||||
elif mask.dim() == 3:
|
||||
# [batch, q_len, kv_len] → [batch, 1, q_len, kv_len]
|
||||
attn_mask = mask[:, None, :, :]
|
||||
else:
|
||||
raise ValueError(f"mask must be 2D or 3D, got {mask.dim()}D")
|
||||
|
||||
if causal_offset >= 0:
|
||||
batch = q.size(0)
|
||||
# q row i attends to kv cols 0..(causal_offset + i)
|
||||
q_idx = torch.arange(q_len, device=q.device).unsqueeze(1) # [q_len, 1]
|
||||
kv_idx = torch.arange(kv_len, device=q.device).unsqueeze(0) # [1, kv_len]
|
||||
causal_bool = kv_idx > (causal_offset + q_idx) # True = masked out
|
||||
causal_mask = causal_bool.unsqueeze(0).expand(
|
||||
batch, -1, -1
|
||||
) # [batch, q_len, kv_len]
|
||||
causal_mask = causal_mask[:, None, :, :] # [batch, 1, q_len, kv_len]
|
||||
|
||||
if attn_mask is not None:
|
||||
attn_mask = attn_mask | causal_mask
|
||||
else:
|
||||
attn_mask = causal_mask
|
||||
|
||||
return attn_mask, resolved_scale
|
||||
|
||||
|
||||
def _torch_fallback(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
mask: torch.Tensor | None,
|
||||
causal_offset: int,
|
||||
scale: float,
|
||||
q_layout: int,
|
||||
kv_layout: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Reference attention via ``scaled_dot_product_attention``.
|
||||
|
||||
q_layout / kv_layout: 0 = b h l d, 1 = b l h d.
|
||||
If kv_layout is None, uses q_layout (Q and K/V share the same layout).
|
||||
"""
|
||||
if kv_layout is None:
|
||||
kv_layout = q_layout
|
||||
q = _to_bhld(q, q_layout)
|
||||
k = _to_bhld(k, kv_layout)
|
||||
v = _to_bhld(v, kv_layout)
|
||||
k, v = _expand_kv_heads(k, v, q.size(1))
|
||||
attn_mask, resolved_scale = _build_attn_mask(q, k, mask, causal_offset, scale)
|
||||
out = F.scaled_dot_product_attention(
|
||||
q, k, v, attn_mask=attn_mask, is_causal=False, scale=resolved_scale
|
||||
)
|
||||
# Restore Q's original layout
|
||||
if q_layout == 1:
|
||||
out = out.transpose(1, 2)
|
||||
return out
|
||||
|
||||
|
||||
def _gather_kv_from_pages(
|
||||
page_table: torch.Tensor,
|
||||
k_cache: torch.Tensor,
|
||||
v_cache: torch.Tensor,
|
||||
page_size: int,
|
||||
kv_len: int,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Gather contiguous K/V from paged cache for torch SDPA fallback.
|
||||
|
||||
Shapes:
|
||||
page_table : [batch, max_pages] (int64)
|
||||
k_cache : [n_pages, page_size, n_kv_heads, head_dim]
|
||||
v_cache : same as k_cache
|
||||
Returns:
|
||||
k, v : [batch, kv_len, n_kv_heads, head_dim] (b l h d)
|
||||
"""
|
||||
batch, max_pages = page_table.shape
|
||||
_, ps, n_kv_heads, head_dim = k_cache.shape
|
||||
if ps != page_size:
|
||||
raise ValueError(f"k_cache page_size mismatch: {ps} vs {page_size}")
|
||||
|
||||
# Vectorized gather: build physical page + offset indices, then advanced-index
|
||||
positions = torch.arange(kv_len, device=page_table.device)
|
||||
logical_pages = positions // page_size # [kv_len]
|
||||
page_offsets = positions % page_size # [kv_len]
|
||||
|
||||
phys_pages = page_table[:, logical_pages] # [batch, kv_len]
|
||||
# k_cache[phys_pages, page_offsets] → [batch, kv_len, n_kv_heads, head_dim] (b l h d)
|
||||
k = k_cache[phys_pages, page_offsets]
|
||||
v = v_cache[phys_pages, page_offsets]
|
||||
return k, v
|
||||
|
||||
|
||||
def attn_decode(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
mask: torch.Tensor | None = None,
|
||||
causal_offset: int = -1,
|
||||
scale: float = 0.0,
|
||||
layout: str = "bhld",
|
||||
) -> torch.Tensor:
|
||||
li = _parse_layout(layout)
|
||||
if _available["attn_decode"]:
|
||||
return _modules["attn_decode"].attn_decode(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
mask=mask,
|
||||
causal_offset=causal_offset,
|
||||
scale=scale,
|
||||
layout=li,
|
||||
)
|
||||
return _torch_fallback(q, k, v, mask, causal_offset, scale, q_layout=li)
|
||||
|
||||
|
||||
def attn_prefill(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
mask: torch.Tensor | None = None,
|
||||
causal_offset: int = -1,
|
||||
scale: float = 0.0,
|
||||
layout: str = "bhld",
|
||||
) -> torch.Tensor:
|
||||
li = _parse_layout(layout)
|
||||
if _available["attn_prefill"]:
|
||||
return _modules["attn_prefill"].attn_prefill(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
mask=mask,
|
||||
causal_offset=causal_offset,
|
||||
scale=scale,
|
||||
layout=li,
|
||||
)
|
||||
return _torch_fallback(q, k, v, mask, causal_offset, scale, q_layout=li)
|
||||
|
||||
|
||||
def attn_paged_decode(
|
||||
q: torch.Tensor,
|
||||
page_table: torch.Tensor,
|
||||
k_cache: torch.Tensor,
|
||||
v_cache: torch.Tensor,
|
||||
page_size: int,
|
||||
kv_len: int,
|
||||
mask: torch.Tensor | None = None,
|
||||
causal_offset: int = -1,
|
||||
scale: float = 0.0,
|
||||
layout: str = "bhld",
|
||||
) -> torch.Tensor:
|
||||
li = _parse_layout(layout)
|
||||
if _available["attn_paged_decode"]:
|
||||
return _modules["attn_paged_decode"].attn_paged_decode(
|
||||
q,
|
||||
page_table,
|
||||
k_cache,
|
||||
v_cache,
|
||||
page_size,
|
||||
kv_len,
|
||||
mask=mask,
|
||||
causal_offset=causal_offset,
|
||||
scale=scale,
|
||||
layout=li,
|
||||
)
|
||||
# Gathered K/V are always b l h d
|
||||
k, v = _gather_kv_from_pages(page_table, k_cache, v_cache, page_size, kv_len)
|
||||
return _torch_fallback(
|
||||
q, k, v, mask, causal_offset, scale, q_layout=li, kv_layout=1
|
||||
)
|
||||
+1
-2
@@ -4,7 +4,6 @@ import inspect
|
||||
import sys
|
||||
from abc import ABC
|
||||
from typing import (
|
||||
Any,
|
||||
Callable,
|
||||
Dict,
|
||||
ForwardRef,
|
||||
@@ -38,7 +37,7 @@ def _resolve_type(
|
||||
ns = vars(mod)
|
||||
|
||||
if isinstance(arg, ForwardRef):
|
||||
return arg._evaluate(ns, None, frozenset(), recursive_guard=frozenset())
|
||||
return arg._evaluate(ns, None, recursive_guard=frozenset())
|
||||
|
||||
return ns.get(name)
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@ Layers:
|
||||
- protocols/: Response builders (OpenAI, Anthropic)
|
||||
- transport/: SSE transport utilities
|
||||
- engine.py: Facade (InferenceEngine), Value Object (GenerationRequest)
|
||||
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy)
|
||||
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy, FrequencyPenaltyStrategy)
|
||||
"""
|
||||
|
||||
from astrai.inference.api import (
|
||||
@@ -30,10 +30,14 @@ from astrai.inference.api.openai import OpenAIResponseBuilder
|
||||
from astrai.inference.core import (
|
||||
STOP,
|
||||
Allocator,
|
||||
CacheView,
|
||||
ContiguousCache,
|
||||
ContiguousCacheView,
|
||||
Executor,
|
||||
InferenceScheduler,
|
||||
KVCache,
|
||||
KvcacheView,
|
||||
PageCache,
|
||||
PageCacheView,
|
||||
PagePool,
|
||||
PrefixCache,
|
||||
Storage,
|
||||
@@ -46,6 +50,7 @@ from astrai.inference.core import (
|
||||
from astrai.inference.engine import GenerationRequest, InferenceEngine
|
||||
from astrai.inference.sample import (
|
||||
BaseSamplingStrategy,
|
||||
FrequencyPenaltyStrategy,
|
||||
SamplingPipeline,
|
||||
TemperatureStrategy,
|
||||
TopKStrategy,
|
||||
@@ -63,8 +68,12 @@ __all__ = [
|
||||
"TaskManager",
|
||||
"TaskStatus",
|
||||
"Allocator",
|
||||
"CacheView",
|
||||
"KVCache",
|
||||
"KvcacheView",
|
||||
"ContiguousCache",
|
||||
"ContiguousCacheView",
|
||||
"PageCache",
|
||||
"PageCacheView",
|
||||
"PagePool",
|
||||
"PrefixCache",
|
||||
"Storage",
|
||||
@@ -75,6 +84,7 @@ __all__ = [
|
||||
"TemperatureStrategy",
|
||||
"TopKStrategy",
|
||||
"TopPStrategy",
|
||||
"FrequencyPenaltyStrategy",
|
||||
"SamplingPipeline",
|
||||
"ProtocolHandler",
|
||||
"StopChecker",
|
||||
|
||||
@@ -21,7 +21,6 @@ logger = logging.getLogger(__name__)
|
||||
_UNSUPPORTED_PARAMS = (
|
||||
"n",
|
||||
"presence_penalty",
|
||||
"frequency_penalty",
|
||||
"logit_bias",
|
||||
"user",
|
||||
)
|
||||
|
||||
@@ -125,6 +125,7 @@ class ProtocolHandler:
|
||||
temperature=self.request.temperature,
|
||||
top_p=self.request.top_p,
|
||||
top_k=self.request.top_k,
|
||||
frequency_penalty=getattr(self.request, "frequency_penalty", 0.0),
|
||||
)
|
||||
|
||||
if self.request.stream:
|
||||
|
||||
@@ -7,6 +7,7 @@ Subclasses may optionally consume ``token_ids`` for token-level parsing
|
||||
(e.g. Harmony / VLM-style parsers).
|
||||
"""
|
||||
|
||||
import json
|
||||
import re
|
||||
import uuid
|
||||
from abc import ABC, abstractmethod
|
||||
@@ -117,6 +118,29 @@ def _parse_tool_call_json(json_str: str, complete: bool):
|
||||
|
||||
Returns ``(name, args, valid)``.
|
||||
"""
|
||||
if complete:
|
||||
try:
|
||||
obj = json.loads(json_str)
|
||||
except json.JSONDecodeError:
|
||||
return None, "", False
|
||||
name = obj.get("name")
|
||||
if not isinstance(name, str) or not name:
|
||||
return None, "", False
|
||||
args = obj.get("arguments")
|
||||
if isinstance(args, dict):
|
||||
if not args:
|
||||
args = ""
|
||||
else:
|
||||
args = json.dumps(args, ensure_ascii=False)
|
||||
args = args[1:-1].rstrip()
|
||||
elif isinstance(args, list):
|
||||
args = json.dumps(args, ensure_ascii=False) if args else ""
|
||||
elif isinstance(args, str):
|
||||
pass
|
||||
else:
|
||||
args = str(args) if args is not None else ""
|
||||
return name, args, True
|
||||
|
||||
name_match = re.search(r'"name"\s*:\s*"([^"]*)"', json_str)
|
||||
if not name_match:
|
||||
return None, "", False
|
||||
@@ -127,8 +151,6 @@ def _parse_tool_call_json(json_str: str, complete: bool):
|
||||
return name, "", True
|
||||
|
||||
raw = args_match.group(1).rstrip()
|
||||
if complete and raw.endswith("}"):
|
||||
raw = raw[:-1].rstrip()
|
||||
if raw.startswith("{"):
|
||||
inner = raw[1:].rstrip()
|
||||
if inner.endswith("}"):
|
||||
@@ -156,9 +178,6 @@ def _find_tool_calls(text: str, start_pos: int = 0):
|
||||
break
|
||||
|
||||
json_str = text[brace:end]
|
||||
if not _TOOL_CALL_HEAD_RE.search(json_str):
|
||||
pos = end
|
||||
continue
|
||||
|
||||
name, args, valid = _parse_tool_call_json(json_str, complete=True)
|
||||
if not valid or name is None:
|
||||
@@ -186,7 +205,7 @@ def _find_partial_tool_call(text: str, start_pos: int = 0):
|
||||
return None
|
||||
|
||||
json_str = text[brace:]
|
||||
if not _TOOL_CALL_HEAD_RE.search(json_str):
|
||||
if '"name"' not in json_str:
|
||||
return None
|
||||
|
||||
name, args, valid = _parse_tool_call_json(json_str, complete=False)
|
||||
|
||||
@@ -2,8 +2,12 @@
|
||||
|
||||
from astrai.inference.core.cache import (
|
||||
Allocator,
|
||||
CacheView,
|
||||
ContiguousCache,
|
||||
ContiguousCacheView,
|
||||
KVCache,
|
||||
KvcacheView,
|
||||
PageCache,
|
||||
PageCacheView,
|
||||
PagePool,
|
||||
PrefixCache,
|
||||
Storage,
|
||||
@@ -16,8 +20,12 @@ from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
|
||||
|
||||
__all__ = [
|
||||
"Allocator",
|
||||
"CacheView",
|
||||
"KVCache",
|
||||
"KvcacheView",
|
||||
"ContiguousCache",
|
||||
"ContiguousCacheView",
|
||||
"PageCache",
|
||||
"PageCacheView",
|
||||
"PagePool",
|
||||
"PrefixCache",
|
||||
"Storage",
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import threading
|
||||
from abc import ABC, abstractmethod
|
||||
from collections import OrderedDict
|
||||
from typing import Callable, Dict, List, Optional, Tuple
|
||||
|
||||
@@ -62,7 +63,8 @@ class Allocator:
|
||||
|
||||
def touch(self, idx: int):
|
||||
with self._lock:
|
||||
self._lru.move_to_end(idx)
|
||||
if idx in self._lru:
|
||||
self._lru.move_to_end(idx)
|
||||
|
||||
|
||||
class PrefixCache:
|
||||
@@ -274,7 +276,46 @@ class Storage:
|
||||
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,
|
||||
write_positions: Optional[Tensor] = None,
|
||||
) -> CacheView: ...
|
||||
|
||||
def task_cached(self, task_id: str) -> int:
|
||||
return 0
|
||||
|
||||
def task_record_hashes(
|
||||
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
||||
): ...
|
||||
|
||||
|
||||
class PageCacheView(CacheView):
|
||||
"""Bundles Storage + page_table + total_len for attention layers."""
|
||||
|
||||
def __init__(self, storage: Storage, page_table: Tensor, total_len: int = 0):
|
||||
@@ -290,8 +331,8 @@ class KvcacheView:
|
||||
return self._storage.gather(layer_id, self._page_table, self._total_len)
|
||||
|
||||
|
||||
class KVCache:
|
||||
"""Facade: page management + KV-cache I/O for continuous batching."""
|
||||
class PageCache(KVCache):
|
||||
"""Paged KV-cache with prefix sharing."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -361,8 +402,132 @@ class KVCache:
|
||||
for i in range(start_logical_page, full_pages):
|
||||
self._pool.record(page_table[i], prompt_ids, i)
|
||||
|
||||
def make_table_tensor(self, task_ids: List[str], device: torch.device) -> Tensor:
|
||||
return self._table.table_tensor(task_ids, device)
|
||||
def bind_tasks(
|
||||
self,
|
||||
task_ids: List[str],
|
||||
total_len: int,
|
||||
device: torch.device,
|
||||
write_positions: Optional[Tensor] = None,
|
||||
) -> PageCacheView:
|
||||
page_table = self._table.table_tensor(task_ids, device)
|
||||
return PageCacheView(self._storage, page_table, total_len)
|
||||
|
||||
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,
|
||||
write_positions: Optional[Tensor] = None,
|
||||
):
|
||||
self._cache = cache
|
||||
self._batch_indices = batch_indices
|
||||
self._total_len = total_len
|
||||
self._write_positions = write_positions
|
||||
|
||||
def write(self, layer_id: int, k: Tensor, v: Tensor):
|
||||
seq_len = k.size(1)
|
||||
indices = self._batch_indices
|
||||
if self._write_positions is not None and seq_len == 1:
|
||||
pos = self._write_positions
|
||||
self._cache.k[layer_id, indices, pos] = k.squeeze(1)
|
||||
self._cache.v[layer_id, indices, pos] = v.squeeze(1)
|
||||
for s, p in zip(indices.tolist(), pos.tolist()):
|
||||
cur = self._cache._slot_len.get(s, 0)
|
||||
if p + 1 > cur:
|
||||
self._cache._slot_len[s] = p + 1
|
||||
else:
|
||||
start_pos = self._total_len - seq_len
|
||||
self._cache.k[layer_id, indices, start_pos : start_pos + seq_len] = k
|
||||
self._cache.v[layer_id, indices, start_pos : start_pos + seq_len] = v
|
||||
new_len = start_pos + seq_len
|
||||
for s in indices.tolist():
|
||||
cur = self._cache._slot_len.get(s, 0)
|
||||
if new_len > cur:
|
||||
self._cache._slot_len[s] = new_len
|
||||
|
||||
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
|
||||
max_len = max(
|
||||
self._cache._slot_len.get(int(s), 0) for s in self._batch_indices.tolist()
|
||||
)
|
||||
indices = self._batch_indices
|
||||
k = self._cache.k[layer_id, indices, :max_len]
|
||||
v = self._cache.v[layer_id, indices, :max_len]
|
||||
return k, v
|
||||
|
||||
|
||||
class ContiguousCache(KVCache):
|
||||
"""Contiguous per-slot KV cache (default implementation)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
n_layers: int,
|
||||
max_batch_size: int,
|
||||
max_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 task_cached(self, task_id: str) -> int:
|
||||
slot = self._task_slot.get(task_id)
|
||||
if slot is None:
|
||||
return 0
|
||||
return self._slot_len.get(slot, 0)
|
||||
|
||||
def bind_tasks(
|
||||
self,
|
||||
task_ids: List[str],
|
||||
total_len: int,
|
||||
device: torch.device,
|
||||
write_positions: Optional[Tensor] = None,
|
||||
) -> ContiguousCacheView:
|
||||
slots = [self._task_slot[tid] for tid in task_ids]
|
||||
batch_indices = torch.tensor(slots, dtype=torch.long, device=device)
|
||||
return ContiguousCacheView(
|
||||
self, batch_indices, total_len, write_positions=write_positions
|
||||
)
|
||||
|
||||
@@ -19,13 +19,13 @@ class Executor:
|
||||
self,
|
||||
model: AutoModel,
|
||||
tokenizer: AutoTokenizer,
|
||||
page_cache: KVCache,
|
||||
kv_cache: KVCache,
|
||||
device: Optional[str] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
self.model = model
|
||||
self.tokenizer = tokenizer
|
||||
self.page_cache = page_cache
|
||||
self.kv_cache = kv_cache
|
||||
self.device = device or next(model.parameters()).device
|
||||
self.dtype = dtype or next(model.parameters()).dtype
|
||||
|
||||
@@ -43,7 +43,6 @@ class Executor:
|
||||
)
|
||||
|
||||
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():
|
||||
self.model(
|
||||
@@ -53,7 +52,7 @@ class Executor:
|
||||
)
|
||||
.unsqueeze(0)
|
||||
.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]:
|
||||
@@ -72,16 +71,47 @@ class Executor:
|
||||
total_len = position_ids.max().item() + 1
|
||||
|
||||
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)
|
||||
top_ks = torch.tensor([t.top_k for t in tasks], device=self.device)
|
||||
top_ps = torch.tensor([t.top_p for t in tasks], device=self.device)
|
||||
freq_penalties = torch.tensor(
|
||||
[t.frequency_penalty for t in tasks], device=self.device
|
||||
)
|
||||
|
||||
history_lists = []
|
||||
mask_lists = []
|
||||
for t in tasks:
|
||||
window = t.rep_window
|
||||
prompt_part = t.prompt_ids[-window:]
|
||||
ids = prompt_part + t.output_ids
|
||||
history_lists.append(ids)
|
||||
mask_lists.append([True] * len(ids))
|
||||
|
||||
max_len = max(len(h) for h in history_lists)
|
||||
padded_ids = torch.zeros(
|
||||
len(tasks), max_len, dtype=torch.long, device=self.device
|
||||
)
|
||||
padded_mask = torch.zeros(
|
||||
len(tasks), max_len, dtype=torch.bool, device=self.device
|
||||
)
|
||||
for i, (h, m) in enumerate(zip(history_lists, mask_lists)):
|
||||
padded_ids[i, : len(h)] = torch.tensor(
|
||||
h, dtype=torch.long, device=self.device
|
||||
)
|
||||
padded_mask[i, : len(m)] = torch.tensor(
|
||||
m, dtype=torch.bool, device=self.device
|
||||
)
|
||||
|
||||
with torch.inference_mode():
|
||||
outputs = self.model(
|
||||
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,
|
||||
write_positions=position_ids,
|
||||
),
|
||||
position_ids=position_ids.unsqueeze(1),
|
||||
)
|
||||
logits = outputs["logits"][:, -1, :]
|
||||
@@ -91,4 +121,7 @@ class Executor:
|
||||
temperature=temperatures,
|
||||
top_k=top_ks,
|
||||
top_p=top_ps,
|
||||
frequency_penalty=freq_penalties,
|
||||
input_ids=padded_ids,
|
||||
input_mask=padded_mask,
|
||||
).tolist()
|
||||
|
||||
@@ -4,7 +4,7 @@ from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
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.task import STOP, Task, TaskManager, TaskStatus
|
||||
from astrai.model.automodel import AutoModel
|
||||
@@ -14,7 +14,7 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class InferenceScheduler:
|
||||
"""Four-phase continuous batching loop: cleanup -> refill -> prefill -> decode."""
|
||||
"""Continuous batching loop: cleanup -> refill -> prefill -> decode (all groups)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -23,9 +23,9 @@ class InferenceScheduler:
|
||||
max_batch_size: int = 16,
|
||||
max_seq_len: Optional[int] = None,
|
||||
max_prompt_len: int = 2048,
|
||||
page_size: int = 64,
|
||||
device: Optional[str] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
cache: Optional[KVCache] = None,
|
||||
):
|
||||
config = model.config
|
||||
|
||||
@@ -41,19 +41,20 @@ class InferenceScheduler:
|
||||
self.device = device or next(model.parameters()).device
|
||||
self.dtype = dtype or next(model.parameters()).dtype
|
||||
|
||||
n_pages = (
|
||||
max_batch_size * (self.max_seq_len + page_size) + page_size - 1
|
||||
) // page_size
|
||||
head_dim = config.dim // config.n_heads
|
||||
|
||||
self._page_cache = KVCache(
|
||||
config.n_layers,
|
||||
n_pages,
|
||||
page_size,
|
||||
config.n_kv_heads,
|
||||
config.dim // config.n_heads,
|
||||
self.device,
|
||||
self.dtype,
|
||||
)
|
||||
if cache is not None:
|
||||
self._cache = cache
|
||||
else:
|
||||
self._cache = ContiguousCache(
|
||||
config.n_layers,
|
||||
max_batch_size,
|
||||
self.max_seq_len,
|
||||
config.n_kv_heads,
|
||||
head_dim,
|
||||
self.device,
|
||||
self.dtype,
|
||||
)
|
||||
|
||||
self._task_mgr = TaskManager(
|
||||
tokenizer=tokenizer,
|
||||
@@ -65,7 +66,7 @@ class InferenceScheduler:
|
||||
self._executor = Executor(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
page_cache=self._page_cache,
|
||||
kv_cache=self._cache,
|
||||
device=self.device,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
@@ -78,18 +79,19 @@ class InferenceScheduler:
|
||||
|
||||
def remove_task(self, task_id: str):
|
||||
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]:
|
||||
return self._task_mgr.get_stats()
|
||||
|
||||
def _run_generation_loop(self):
|
||||
stop_ids = self._task_mgr.tokenizer.stop_ids
|
||||
cache = self._cache
|
||||
try:
|
||||
while not self._stop_event.is_set():
|
||||
finished = self._task_mgr.remove_finished_tasks(stop_ids)
|
||||
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()
|
||||
available = self._task_mgr.max_batch_size - len(active)
|
||||
@@ -97,7 +99,7 @@ class InferenceScheduler:
|
||||
candidates = self._task_mgr.pull_candidates(available)
|
||||
failed = []
|
||||
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)
|
||||
else:
|
||||
failed.append(task)
|
||||
@@ -112,7 +114,7 @@ class InferenceScheduler:
|
||||
t
|
||||
for t in self._task_mgr.get_active_tasks()
|
||||
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:
|
||||
for t in to_prefill:
|
||||
@@ -122,69 +124,55 @@ class InferenceScheduler:
|
||||
for t in to_prefill:
|
||||
key = (
|
||||
len(t.prompt_ids),
|
||||
self._page_cache.task_cached(t.task_id),
|
||||
cache.task_cached(t.task_id),
|
||||
)
|
||||
groups.setdefault(key, []).append(t)
|
||||
|
||||
for (prompt_len, start_pos), group in groups.items():
|
||||
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:
|
||||
self._page_cache.task_record_hashes(
|
||||
t.task_id,
|
||||
t.prompt_ids,
|
||||
start_logical_page=start_logical_page,
|
||||
cache.task_record_hashes(
|
||||
t.task_id, t.prompt_ids, start_logical_page
|
||||
)
|
||||
|
||||
pos_groups: Dict[int, List[Task]] = {}
|
||||
for t in self._task_mgr.get_active_tasks():
|
||||
pos_groups.setdefault(t.next_pos, []).append(t)
|
||||
decode_tasks = self._task_mgr.get_active_tasks()
|
||||
|
||||
if pos_groups:
|
||||
best_key = max(pos_groups, key=lambda k: len(pos_groups[k]))
|
||||
group = sorted(pos_groups[best_key], key=lambda t: t.task_id)
|
||||
valid: List[Task] = []
|
||||
for t in sorted(decode_tasks, key=lambda t: t.task_id):
|
||||
if cache.task_extend(t.task_id, t.next_pos):
|
||||
valid.append(t)
|
||||
else:
|
||||
t.status = TaskStatus.ABORTED
|
||||
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||
|
||||
valid: List[Task] = []
|
||||
for t in group:
|
||||
if self._page_cache.task_extend(t.task_id, t.next_pos):
|
||||
valid.append(t)
|
||||
else:
|
||||
t.status = TaskStatus.ABORTED
|
||||
if t.stream_callback:
|
||||
t.stream_callback(STOP)
|
||||
if valid:
|
||||
next_tokens = self._executor.execute_decode(valid)
|
||||
|
||||
if valid:
|
||||
next_tokens = self._executor.execute_decode(valid)
|
||||
for t, ntok in zip(valid, next_tokens):
|
||||
t.output_ids.append(ntok)
|
||||
t.output_tokens += 1
|
||||
new_text = t.decode_new_token(self._task_mgr.tokenizer)
|
||||
if new_text:
|
||||
self._task_mgr.invoke_callback(t.task_id, new_text)
|
||||
|
||||
for t, ntok in zip(valid, next_tokens):
|
||||
t.output_ids.append(ntok)
|
||||
t.output_tokens += 1
|
||||
pos = t.input_tokens + t.output_tokens
|
||||
extend_ok = self._page_cache.task_extend(t.task_id, pos)
|
||||
if t.stream_callback:
|
||||
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:
|
||||
if t.is_finished(stop_ids):
|
||||
if t.stream_callback:
|
||||
t.stream_callback(STOP)
|
||||
for t in valid:
|
||||
if t.is_finished(stop_ids):
|
||||
remaining = t.flush_remaining(self._task_mgr.tokenizer)
|
||||
if remaining:
|
||||
self._task_mgr.invoke_callback(t.task_id, remaining)
|
||||
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||
|
||||
except Exception as e:
|
||||
self._stop_event.set()
|
||||
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
|
||||
for task in self._task_mgr.get_active_tasks():
|
||||
if task.stream_callback:
|
||||
task.stream_callback(STOP)
|
||||
self._page_cache.task_free(task.task_id)
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
cache.task_free(task.task_id)
|
||||
for task in self._task_mgr.get_waiting_tasks():
|
||||
if task.stream_callback:
|
||||
task.stream_callback(STOP)
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
self._task_mgr.clear_queues()
|
||||
|
||||
def start(self):
|
||||
@@ -202,12 +190,10 @@ class InferenceScheduler:
|
||||
self._loop_thread.join(timeout=2.0)
|
||||
self._loop_thread = None
|
||||
for task in self._task_mgr.get_active_tasks():
|
||||
if task.stream_callback:
|
||||
task.stream_callback(STOP)
|
||||
self._page_cache.task_free(task.task_id)
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
self._cache.task_free(task.task_id)
|
||||
for task in self._task_mgr.get_waiting_tasks():
|
||||
if task.stream_callback:
|
||||
task.stream_callback(STOP)
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
self._task_mgr.clear_queues()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
@@ -13,6 +13,40 @@ logger = logging.getLogger(__name__)
|
||||
STOP = object()
|
||||
|
||||
|
||||
class StreamDecoder:
|
||||
"""Incremental decoder for byte-level BPE streaming.
|
||||
|
||||
Byte-level BPE may split a single Unicode character (e.g. em-dash,
|
||||
smart quotes) across multiple tokens. Decoding such a token in
|
||||
isolation produces U+FFFD (replacement char). This decoder
|
||||
accumulates token IDs and only emits text once the trailing
|
||||
characters are complete, buffering incomplete multi-byte sequences
|
||||
until the next token arrives.
|
||||
"""
|
||||
|
||||
__slots__ = ("_tokenizer", "_ids", "_emitted")
|
||||
|
||||
def __init__(self, tokenizer: AutoTokenizer):
|
||||
self._tokenizer = tokenizer
|
||||
self._ids: List[int] = []
|
||||
self._emitted: str = ""
|
||||
|
||||
def push(self, token_id: int) -> str:
|
||||
"""Append a token ID and return newly completed text.
|
||||
|
||||
Returns "" while a multi-byte character is still incomplete.
|
||||
"""
|
||||
self._ids.append(token_id)
|
||||
full = self._tokenizer.decode(self._ids, skip_special_tokens=True)
|
||||
if full.endswith("\ufffd"):
|
||||
return ""
|
||||
if len(full) > len(self._emitted):
|
||||
diff = full[len(self._emitted) :]
|
||||
self._emitted = full
|
||||
return diff
|
||||
return ""
|
||||
|
||||
|
||||
class TaskStatus(Enum):
|
||||
"""Task lifecycle states."""
|
||||
|
||||
@@ -33,7 +67,8 @@ class Task:
|
||||
temperature: float = 1.0,
|
||||
top_p: float = 1.0,
|
||||
top_k: int = 50,
|
||||
stream_callback: Optional[Callable[[str], None]] = None,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
):
|
||||
self.task_id = task_id
|
||||
self.prompt_ids = prompt_ids
|
||||
@@ -41,6 +76,8 @@ class Task:
|
||||
self.temperature = temperature
|
||||
self.top_p = top_p
|
||||
self.top_k = top_k
|
||||
self.frequency_penalty = frequency_penalty
|
||||
self.rep_window = rep_window
|
||||
|
||||
self.status = TaskStatus.PENDING
|
||||
self.output_ids: List[int] = []
|
||||
@@ -48,7 +85,34 @@ class Task:
|
||||
self.output_tokens: int = 0
|
||||
self.arrival_time = time.time()
|
||||
self.finish_time: Optional[float] = None
|
||||
self.stream_callback = stream_callback
|
||||
self._decoder: Optional[StreamDecoder] = None
|
||||
|
||||
def decode_new_token(self, tokenizer: AutoTokenizer) -> str:
|
||||
"""Decode the last appended output token, buffering incomplete
|
||||
multi-byte sequences across calls.
|
||||
|
||||
Lazily creates a :class:`StreamDecoder` on first use.
|
||||
"""
|
||||
if self._decoder is None:
|
||||
self._decoder = StreamDecoder(tokenizer)
|
||||
return self._decoder.push(self.output_ids[-1])
|
||||
|
||||
def flush_remaining(self, tokenizer: AutoTokenizer) -> str:
|
||||
"""Emit any text still buffered in the decoder.
|
||||
|
||||
Called when generation terminates (max_tokens reached, stop
|
||||
sequence, or external removal) to avoid dropping a final
|
||||
incomplete-looking fragment that is actually complete when
|
||||
adjacent to the stop token.
|
||||
"""
|
||||
if self._decoder is None or not self.output_ids:
|
||||
return ""
|
||||
full = tokenizer.decode(self.output_ids, skip_special_tokens=True)
|
||||
if len(full) > len(self._decoder._emitted):
|
||||
diff = full[len(self._decoder._emitted) :]
|
||||
self._decoder._emitted = full
|
||||
return diff
|
||||
return ""
|
||||
|
||||
@property
|
||||
def next_pos(self) -> int:
|
||||
@@ -79,6 +143,7 @@ class TaskManager:
|
||||
|
||||
self.waiting_queue: Deque[Task] = deque()
|
||||
self.active_tasks: List[Task] = []
|
||||
self._callbacks: Dict[str, Callable[[str], None]] = {}
|
||||
|
||||
self._task_event = threading.Event()
|
||||
self._lock = threading.Lock()
|
||||
@@ -93,6 +158,8 @@ class TaskManager:
|
||||
temperature: float = 1.0,
|
||||
top_p: float = 1.0,
|
||||
top_k: int = 50,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
stream_callback: Optional[Callable[[str], None]] = None,
|
||||
) -> str:
|
||||
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
|
||||
@@ -117,12 +184,15 @@ class TaskManager:
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
stream_callback=stream_callback,
|
||||
frequency_penalty=frequency_penalty,
|
||||
rep_window=rep_window,
|
||||
)
|
||||
|
||||
with self._lock:
|
||||
self.waiting_queue.append(task)
|
||||
self._total_tasks += 1
|
||||
if stream_callback:
|
||||
self._callbacks[task_id] = stream_callback
|
||||
|
||||
self._task_event.set()
|
||||
return task_id
|
||||
@@ -134,8 +204,14 @@ class TaskManager:
|
||||
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._callbacks.pop(task_id, None)
|
||||
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]:
|
||||
return {
|
||||
"total_tasks": self._total_tasks,
|
||||
@@ -204,6 +280,7 @@ class TaskManager:
|
||||
with self._lock:
|
||||
self.waiting_queue.clear()
|
||||
self.active_tasks.clear()
|
||||
self._callbacks.clear()
|
||||
|
||||
def wake(self):
|
||||
self._task_event.set()
|
||||
|
||||
@@ -8,6 +8,7 @@ from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Tuple,
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from astrai.inference.core.cache import KVCache
|
||||
from astrai.inference.core.scheduler import InferenceScheduler
|
||||
from astrai.inference.core.task import STOP
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
@@ -73,20 +74,31 @@ class GenerationRequest:
|
||||
top_p: float = 1.0,
|
||||
temperature: float = 1.0,
|
||||
max_tokens: Optional[int] = None,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
stream: bool = False,
|
||||
):
|
||||
if not (isinstance(top_k, int) and top_k >= 0):
|
||||
raise ValueError("top_k must be a non-negative integer")
|
||||
if not (0.0 <= top_p <= 1.0):
|
||||
raise ValueError("top_p must be a float between 0.0 and 1.0")
|
||||
if not (isinstance(temperature, (int, float)) and temperature > 0):
|
||||
raise ValueError("temperature must be a positive number")
|
||||
if not (isinstance(temperature, (int, float)) and temperature >= 0):
|
||||
raise ValueError("temperature must be a non-negative number")
|
||||
if not (
|
||||
isinstance(frequency_penalty, (int, float))
|
||||
and -2.0 <= frequency_penalty <= 2.0
|
||||
):
|
||||
raise ValueError("frequency_penalty must be between -2.0 and 2.0")
|
||||
if not (isinstance(rep_window, int) and rep_window > 0):
|
||||
raise ValueError("rep_window must be a positive integer")
|
||||
|
||||
self.messages = messages
|
||||
self.top_k = top_k
|
||||
self.top_p = top_p
|
||||
self.temperature = temperature
|
||||
self.max_tokens = max_tokens
|
||||
self.frequency_penalty = frequency_penalty
|
||||
self.rep_window = rep_window
|
||||
self.stream = stream
|
||||
|
||||
|
||||
@@ -101,6 +113,7 @@ class InferenceEngine:
|
||||
max_seq_len: Optional[int] = None,
|
||||
max_prompt_len: int = 2048,
|
||||
page_size: int = 128,
|
||||
cache: Optional[KVCache] = None,
|
||||
):
|
||||
self.model = model
|
||||
self.tokenizer = tokenizer
|
||||
@@ -110,7 +123,7 @@ class InferenceEngine:
|
||||
max_batch_size=max_batch_size,
|
||||
max_seq_len=max_seq_len,
|
||||
max_prompt_len=max_prompt_len,
|
||||
page_size=page_size,
|
||||
cache=cache,
|
||||
)
|
||||
|
||||
self.scheduler.start()
|
||||
@@ -130,17 +143,33 @@ class InferenceEngine:
|
||||
temperature: float = 1.0,
|
||||
top_p: float = 1.0,
|
||||
top_k: int = 50,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
) -> Union[Generator, str, List[str]]:
|
||||
is_batch = isinstance(prompt, list)
|
||||
prompts = prompt if is_batch else [prompt]
|
||||
|
||||
if stream:
|
||||
return self._generate_streaming(
|
||||
prompts, is_batch, max_tokens, temperature, top_p, top_k
|
||||
prompts,
|
||||
is_batch,
|
||||
max_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
frequency_penalty,
|
||||
rep_window,
|
||||
)
|
||||
else:
|
||||
return self._generate_non_streaming(
|
||||
prompts, is_batch, max_tokens, temperature, top_p, top_k
|
||||
prompts,
|
||||
is_batch,
|
||||
max_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
frequency_penalty,
|
||||
rep_window,
|
||||
)
|
||||
|
||||
def generate_async(
|
||||
@@ -150,9 +179,18 @@ class InferenceEngine:
|
||||
temperature: float = 1.0,
|
||||
top_p: float = 1.0,
|
||||
top_k: int = 50,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
sync_gen = self._generate_streaming(
|
||||
[prompt], False, max_tokens, temperature, top_p, top_k
|
||||
[prompt],
|
||||
False,
|
||||
max_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
frequency_penalty,
|
||||
rep_window,
|
||||
)
|
||||
|
||||
async def _agen():
|
||||
@@ -183,6 +221,8 @@ class InferenceEngine:
|
||||
temperature=request.temperature,
|
||||
top_p=request.top_p,
|
||||
top_k=request.top_k,
|
||||
frequency_penalty=request.frequency_penalty,
|
||||
rep_window=request.rep_window,
|
||||
)
|
||||
|
||||
def _submit_tasks(
|
||||
@@ -192,6 +232,8 @@ class InferenceEngine:
|
||||
temperature: float,
|
||||
top_p: float,
|
||||
top_k: int,
|
||||
frequency_penalty: float,
|
||||
rep_window: int,
|
||||
) -> Tuple[GenerateResult, List[str]]:
|
||||
n = len(prompts)
|
||||
result = GenerateResult(count=n)
|
||||
@@ -204,6 +246,8 @@ class InferenceEngine:
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
frequency_penalty=frequency_penalty,
|
||||
rep_window=rep_window,
|
||||
stream_callback=cb,
|
||||
)
|
||||
task_ids.append(task_id)
|
||||
@@ -224,9 +268,17 @@ class InferenceEngine:
|
||||
temperature: float,
|
||||
top_p: float,
|
||||
top_k: int,
|
||||
frequency_penalty: float,
|
||||
rep_window: int,
|
||||
) -> Generator:
|
||||
result, task_ids = self._submit_tasks(
|
||||
prompts, max_tokens, temperature, top_p, top_k
|
||||
prompts,
|
||||
max_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
frequency_penalty,
|
||||
rep_window,
|
||||
)
|
||||
n = len(prompts)
|
||||
remaining = n
|
||||
@@ -260,9 +312,17 @@ class InferenceEngine:
|
||||
temperature: float,
|
||||
top_p: float,
|
||||
top_k: int,
|
||||
frequency_penalty: float,
|
||||
rep_window: int,
|
||||
) -> Union[str, List[str]]:
|
||||
result, task_ids = self._submit_tasks(
|
||||
prompts, max_tokens, temperature, top_p, top_k
|
||||
prompts,
|
||||
max_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
frequency_penalty,
|
||||
rep_window,
|
||||
)
|
||||
|
||||
try:
|
||||
|
||||
+163
-13
@@ -1,15 +1,15 @@
|
||||
"""Composable sampling strategies for logit transformation.
|
||||
|
||||
Implements the Strategy pattern: each sampling technique
|
||||
(temperature, top-k, top-p) is a pluggable strategy that
|
||||
can be composed into a pipeline.
|
||||
(temperature, top-k, top-p, frequency penalty) is a pluggable
|
||||
strategy that can be composed into a pipeline.
|
||||
|
||||
All strategies accept both scalar and per-sample tensor
|
||||
parameters, so a single pipeline works for any batch size.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List, Union
|
||||
from typing import List, Optional, Union
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
@@ -19,12 +19,23 @@ class BaseSamplingStrategy(ABC):
|
||||
"""Abstract base for a logit transformation strategy."""
|
||||
|
||||
@abstractmethod
|
||||
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
||||
def apply(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
"""Applies the strategy to logits.
|
||||
|
||||
Args:
|
||||
logits: Raw logits tensor (batch, vocab_size).
|
||||
filter_value: Value assigned to filtered-out positions.
|
||||
input_ids: Previously generated token IDs ``[batch, seq_len]``,
|
||||
padded with 0. Used by frequency penalty.
|
||||
input_mask: Boolean mask ``[batch, seq_len]``, True for real
|
||||
tokens, False for padding. Used to exclude padding from
|
||||
penalty computation.
|
||||
|
||||
Returns:
|
||||
Transformed logits tensor.
|
||||
@@ -42,7 +53,13 @@ class TemperatureStrategy(BaseSamplingStrategy):
|
||||
def __init__(self, temperature: Union[float, Tensor] = 1.0):
|
||||
self.temperature = temperature
|
||||
|
||||
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
||||
def apply(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
t = self.temperature
|
||||
if isinstance(t, Tensor):
|
||||
t = t.to(logits.device, non_blocking=True).view(-1, 1)
|
||||
@@ -64,7 +81,13 @@ class TopKStrategy(BaseSamplingStrategy):
|
||||
def __init__(self, top_k: Union[int, Tensor] = 0):
|
||||
self.top_k = top_k
|
||||
|
||||
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
||||
def apply(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
tk = self.top_k
|
||||
if isinstance(tk, Tensor):
|
||||
tk = tk.to(logits.device, non_blocking=True).long().clamp(min=0)
|
||||
@@ -114,7 +137,13 @@ class TopPStrategy(BaseSamplingStrategy):
|
||||
logits[mask] = filter_value
|
||||
return logits
|
||||
|
||||
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
||||
def apply(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
tp = self.top_p
|
||||
if isinstance(tp, Tensor):
|
||||
tp = tp.to(logits.device, non_blocking=True)
|
||||
@@ -125,6 +154,84 @@ class TopPStrategy(BaseSamplingStrategy):
|
||||
return logits
|
||||
|
||||
|
||||
class FrequencyPenaltyStrategy(BaseSamplingStrategy):
|
||||
"""Penalizes tokens based on how many times they appeared in history.
|
||||
|
||||
Subtracts ``penalty * count(token)`` from each token's logit, where
|
||||
``count(token)`` is the number of occurrences in the generation history
|
||||
(prompt + output). A penalty of ``0.0`` disables the strategy.
|
||||
|
||||
Unlike repetition penalty (which only checks *presence*), frequency
|
||||
penalty scales linearly with occurrence count: the first use is
|
||||
penalized once, the third use three times. This allows natural
|
||||
repetition of common words while suppressing degenerate loops.
|
||||
|
||||
Reference: OpenAI API ``frequency_penalty`` parameter.
|
||||
|
||||
Args:
|
||||
penalty: Scalar or ``[batch]`` tensor (0.0 disables, range -2.0~2.0).
|
||||
"""
|
||||
|
||||
def __init__(self, penalty: Union[float, Tensor] = 0.0):
|
||||
self.penalty = penalty
|
||||
|
||||
def apply(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
if input_ids is None:
|
||||
return logits
|
||||
|
||||
p = self.penalty
|
||||
if isinstance(p, Tensor):
|
||||
p = p.to(logits.device, non_blocking=True).view(-1, 1)
|
||||
if (p == 0.0).all():
|
||||
return logits
|
||||
elif p == 0.0:
|
||||
return logits
|
||||
|
||||
input_ids = input_ids.to(logits.device, non_blocking=True)
|
||||
|
||||
if input_mask is not None:
|
||||
input_mask = input_mask.to(logits.device, non_blocking=True)
|
||||
masked_ids = input_ids.clone()
|
||||
masked_ids[~input_mask] = -1
|
||||
else:
|
||||
masked_ids = input_ids
|
||||
|
||||
batch_sz, seq_len = masked_ids.shape
|
||||
vocab_size = logits.size(-1)
|
||||
|
||||
if isinstance(p, Tensor):
|
||||
penalty_per_row = p.expand(batch_sz, 1)
|
||||
else:
|
||||
penalty_per_row = torch.full(
|
||||
(batch_sz, 1), float(p), device=logits.device, dtype=logits.dtype
|
||||
)
|
||||
|
||||
counts = torch.zeros(
|
||||
batch_sz, vocab_size, device=logits.device, dtype=logits.dtype
|
||||
)
|
||||
valid_mask = masked_ids >= 0
|
||||
if valid_mask.any():
|
||||
valid_ids = masked_ids[valid_mask]
|
||||
row_indices = (
|
||||
torch.arange(batch_sz, device=logits.device)
|
||||
.unsqueeze(1)
|
||||
.expand_as(masked_ids)[valid_mask]
|
||||
)
|
||||
counts.index_put_(
|
||||
(row_indices, valid_ids),
|
||||
torch.ones_like(valid_ids, dtype=logits.dtype),
|
||||
accumulate=True,
|
||||
)
|
||||
|
||||
return logits - penalty_per_row * counts
|
||||
|
||||
|
||||
class SamplingPipeline(BaseSamplingStrategy):
|
||||
"""Composes multiple sampling strategies into a single transformation.
|
||||
|
||||
@@ -145,23 +252,53 @@ class SamplingPipeline(BaseSamplingStrategy):
|
||||
def __init__(self, strategies: List[BaseSamplingStrategy]):
|
||||
self.strategies = strategies
|
||||
|
||||
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
||||
def apply(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
for strategy in self.strategies:
|
||||
logits = strategy.apply(logits, filter_value)
|
||||
logits = strategy.apply(logits, filter_value, input_ids, input_mask)
|
||||
return logits
|
||||
|
||||
@torch.no_grad()
|
||||
def sample(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
||||
@staticmethod
|
||||
def _is_greedy(temperature: Union[float, Tensor]) -> bool:
|
||||
if isinstance(temperature, Tensor):
|
||||
return temperature.numel() == 1 and temperature.item() == 0
|
||||
return temperature == 0
|
||||
|
||||
@torch.inference_mode()
|
||||
def sample(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
"""Apply strategies then sample (softmax + multinomial).
|
||||
|
||||
Short-circuits to ``argmax`` when temperature is exactly 0
|
||||
(deterministic / greedy decode).
|
||||
|
||||
Args:
|
||||
logits: Raw logits ``[batch, vocab_size]``.
|
||||
input_ids: Previously generated token IDs ``[batch, seq_len]``.
|
||||
input_mask: Boolean mask for ``input_ids`` padding.
|
||||
|
||||
Returns:
|
||||
Sampled token IDs ``[batch]``.
|
||||
"""
|
||||
for s in self.strategies:
|
||||
if isinstance(s, TemperatureStrategy) and self._is_greedy(s.temperature):
|
||||
return logits.argmax(dim=-1)
|
||||
break
|
||||
|
||||
return torch.multinomial(
|
||||
torch.softmax(self.apply(logits, filter_value), dim=-1),
|
||||
torch.softmax(
|
||||
self.apply(logits, filter_value, input_ids, input_mask), dim=-1
|
||||
),
|
||||
num_samples=1,
|
||||
).squeeze(-1)
|
||||
|
||||
@@ -172,22 +309,35 @@ def sample(
|
||||
temperature: Union[float, Tensor] = 1.0,
|
||||
top_k: Union[int, Tensor] = 0,
|
||||
top_p: Union[float, Tensor] = 1.0,
|
||||
frequency_penalty: Union[float, Tensor] = 0.0,
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
filter_value: float = -float("inf"),
|
||||
) -> Tensor:
|
||||
"""Apply sampling strategies then sample (softmax + multinomial).
|
||||
|
||||
Shortcut for ``SamplingPipeline(...).sample(logits)``.
|
||||
|
||||
When **temperature** is exactly 0 (scalar or single-element tensor)
|
||||
the function short-circuits to ``argmax`` for deterministic decode.
|
||||
|
||||
Args:
|
||||
logits: Raw logits ``[batch, vocab_size]``.
|
||||
frequency_penalty: Penalty per occurrence for repeated tokens
|
||||
(0.0 disables, range -2.0~2.0).
|
||||
input_ids: Previously generated token IDs ``[batch, seq_len]``.
|
||||
input_mask: Boolean mask for ``input_ids`` padding.
|
||||
|
||||
Returns:
|
||||
Sampled token IDs ``[batch]``.
|
||||
"""
|
||||
if SamplingPipeline._is_greedy(temperature):
|
||||
return logits.argmax(dim=-1)
|
||||
return SamplingPipeline(
|
||||
[
|
||||
TemperatureStrategy(temperature),
|
||||
TopKStrategy(top_k),
|
||||
TopPStrategy(top_p),
|
||||
FrequencyPenaltyStrategy(frequency_penalty),
|
||||
]
|
||||
).sample(logits, filter_value)
|
||||
).sample(logits, filter_value, input_ids, input_mask)
|
||||
|
||||
@@ -6,7 +6,7 @@ import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
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.norm import RMSNorm
|
||||
from astrai.model.components.rope import apply_rotary_emb
|
||||
@@ -38,6 +38,7 @@ class GQA(nn.Module):
|
||||
norm_eps: float,
|
||||
use_gated_attention: bool,
|
||||
layer_id: int,
|
||||
n_layers: int = 1,
|
||||
):
|
||||
super().__init__()
|
||||
assert dim % n_heads == 0
|
||||
@@ -55,7 +56,7 @@ class GQA(nn.Module):
|
||||
self.q_proj = Linear(dim, n_heads * self.head_dim)
|
||||
self.k_proj = Linear(dim, n_kv_heads * self.head_dim)
|
||||
self.v_proj = Linear(dim, n_kv_heads * self.head_dim)
|
||||
self.o_proj = Linear(dim, dim)
|
||||
self.o_proj = Linear(dim, dim, init_std=0.02 / (2 * n_layers) ** 0.5)
|
||||
|
||||
if self.use_qk_norm:
|
||||
self.q_norm = RMSNorm(self.head_dim, norm_eps)
|
||||
@@ -74,7 +75,7 @@ class GQA(nn.Module):
|
||||
x: Tensor,
|
||||
rotary_emb: Tensor,
|
||||
attn_mask: Tensor = None,
|
||||
paged_cache: Optional[KvcacheView] = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
) -> Tensor:
|
||||
is_causal = attn_mask is None
|
||||
|
||||
@@ -121,6 +122,7 @@ class MLA(nn.Module):
|
||||
use_qk_norm: bool,
|
||||
use_gated_attention: bool,
|
||||
layer_id: int,
|
||||
n_layers: int = 1,
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
@@ -148,7 +150,9 @@ class MLA(nn.Module):
|
||||
n_kv_heads * (2 * self.head_dim),
|
||||
)
|
||||
|
||||
self.o_proj = Linear(dim, dim, bias=False)
|
||||
self.o_proj = Linear(
|
||||
dim, dim, bias=False, init_std=0.02 / (2 * n_layers) ** 0.5
|
||||
)
|
||||
|
||||
if use_gated_attention:
|
||||
self.gate = Linear(dim, dim, bias=False)
|
||||
@@ -158,7 +162,7 @@ class MLA(nn.Module):
|
||||
x: Tensor,
|
||||
rotary_emb: Tensor,
|
||||
attn_mask: Tensor = None,
|
||||
paged_cache: Optional[KvcacheView] = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
) -> Tensor:
|
||||
bsz, seq_len, _ = x.size()
|
||||
is_causal = attn_mask is None
|
||||
|
||||
@@ -1,51 +1,31 @@
|
||||
from dataclasses import asdict
|
||||
from typing import Optional
|
||||
|
||||
import torch.nn as nn
|
||||
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.mlp import FFNFactory
|
||||
from astrai.model.components.norm import RMSNorm
|
||||
|
||||
|
||||
class DecoderBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
n_heads: int,
|
||||
dim_ffn: int,
|
||||
n_kv_heads: int,
|
||||
norm_eps: float,
|
||||
use_qk_norm: bool,
|
||||
use_gated_attention: bool,
|
||||
layer_id: int,
|
||||
attn_type: str = "gqa",
|
||||
ffn_type: str = "mlp",
|
||||
**kwargs,
|
||||
):
|
||||
def __init__(self, config, layer_id: int):
|
||||
super().__init__()
|
||||
self.attention = AttnFactory.create(
|
||||
attn_type,
|
||||
dim=dim,
|
||||
n_heads=n_heads,
|
||||
n_kv_heads=n_kv_heads,
|
||||
use_qk_norm=use_qk_norm,
|
||||
norm_eps=norm_eps,
|
||||
use_gated_attention=use_gated_attention,
|
||||
layer_id=layer_id,
|
||||
**kwargs,
|
||||
)
|
||||
self.input_norm = RMSNorm(dim, norm_eps)
|
||||
self.post_attention_norm = RMSNorm(dim, norm_eps)
|
||||
self.mlp = FFNFactory.create(ffn_type, dim, dim_ffn, **kwargs)
|
||||
cfg = asdict(config)
|
||||
cfg["down_init_std"] = 0.02 / (2 * config.n_layers) ** 0.5
|
||||
self.attention = AttnFactory.create(config.attn_type, **cfg, layer_id=layer_id)
|
||||
self.input_norm = RMSNorm(config.dim, config.norm_eps)
|
||||
self.post_attention_norm = RMSNorm(config.dim, config.norm_eps)
|
||||
self.mlp = FFNFactory.create(config.ffn_type, **cfg)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Tensor,
|
||||
rotary_emb: Tensor,
|
||||
attention_mask: Optional[Tensor] = None,
|
||||
paged_cache: Optional[KvcacheView] = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
) -> Tensor:
|
||||
attn_output = self.attention(
|
||||
self.input_norm(x),
|
||||
|
||||
@@ -7,10 +7,13 @@ from torch import Tensor
|
||||
|
||||
|
||||
class Embedding(nn.Module):
|
||||
def __init__(self, vocab_size: int, embedding_dim: int):
|
||||
def __init__(self, vocab_size: int, embedding_dim: int, neftune_alpha: float = 0.0):
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.empty((vocab_size, embedding_dim)))
|
||||
self.neftune_noise_alpha = 0.0
|
||||
self.neftune_noise_alpha = neftune_alpha
|
||||
|
||||
def set_neftune_alpha(self, alpha: float):
|
||||
self.neftune_noise_alpha = alpha
|
||||
|
||||
def reset_parameters(self):
|
||||
nn.init.normal_(self.weight, mean=0.0, std=0.02)
|
||||
|
||||
@@ -5,13 +5,16 @@ from torch import Tensor
|
||||
|
||||
|
||||
class Linear(nn.Module):
|
||||
def __init__(self, in_dim: int, out_dim: int, bias: bool = False):
|
||||
def __init__(
|
||||
self, in_dim: int, out_dim: int, bias: bool = False, init_std: float = 0.02
|
||||
):
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.empty((out_dim, in_dim)))
|
||||
self.bias = nn.Parameter(torch.zeros(out_dim)) if bias else None
|
||||
self.init_std = init_std
|
||||
|
||||
def reset_parameters(self):
|
||||
nn.init.kaiming_uniform_(self.weight, a=5**0.5)
|
||||
nn.init.normal_(self.weight, mean=0.0, std=self.init_std)
|
||||
if self.bias is not None:
|
||||
fan_in, _ = nn.init._calculate_fan_in_and_fan_out(self.weight)
|
||||
bound = 1 / (fan_in**0.5)
|
||||
|
||||
@@ -13,11 +13,11 @@ class FFNFactory(BaseFactory[nn.Module]):
|
||||
|
||||
@FFNFactory.register("mlp")
|
||||
class MLP(nn.Module):
|
||||
def __init__(self, dim: int, dim_ffn: int):
|
||||
def __init__(self, dim: int, dim_ffn: int, down_init_std: float = 0.02):
|
||||
super().__init__()
|
||||
self.up = Linear(dim, dim_ffn)
|
||||
self.gate = Linear(dim, dim_ffn)
|
||||
self.down = Linear(dim_ffn, dim)
|
||||
self.down = Linear(dim_ffn, dim, init_std=down_init_std)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
gated = self.up(x) * F.silu(self.gate(x))
|
||||
@@ -35,6 +35,7 @@ class DeepSeekMoE(nn.Module):
|
||||
n_shared_experts: int = 1,
|
||||
n_activated_experts: int = 2,
|
||||
topk_method: str = "greedy",
|
||||
n_layers: int = 1,
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
@@ -44,12 +45,20 @@ class DeepSeekMoE(nn.Module):
|
||||
self.topk_method = topk_method
|
||||
|
||||
self.router = Linear(dim, n_routed_experts, bias=False)
|
||||
moe_scale = 1 / max(n_shared_experts, 1) + 1 / n_activated_experts
|
||||
down_init_std = 0.02 / (2 * n_layers * moe_scale) ** 0.5
|
||||
|
||||
self.shared_experts = nn.ModuleList(
|
||||
[MLP(dim, dim_ffn) for _ in range(n_shared_experts)]
|
||||
[
|
||||
MLP(dim, dim_ffn, down_init_std=down_init_std)
|
||||
for _ in range(n_shared_experts)
|
||||
]
|
||||
)
|
||||
self.routed_experts = nn.ModuleList(
|
||||
[MLP(dim, dim_ffn) for _ in range(n_routed_experts)]
|
||||
[
|
||||
MLP(dim, dim_ffn, down_init_std=down_init_std)
|
||||
for _ in range(n_routed_experts)
|
||||
]
|
||||
)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
|
||||
+4
-14
@@ -23,22 +23,12 @@ class EmbeddingEncoder(AutoModel):
|
||||
self.rotary_embedding = RotaryEmbedding(
|
||||
rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling
|
||||
)
|
||||
self.embed_tokens = Embedding(config.vocab_size, config.dim)
|
||||
self.embed_tokens = Embedding(
|
||||
config.vocab_size, config.dim, neftune_alpha=config.neftune_alpha
|
||||
)
|
||||
|
||||
self.layers = nn.ModuleList(
|
||||
[
|
||||
DecoderBlock(
|
||||
config.dim,
|
||||
config.n_heads,
|
||||
config.dim_ffn,
|
||||
config.n_kv_heads,
|
||||
config.norm_eps,
|
||||
config.use_qk_norm,
|
||||
config.use_gated_attention,
|
||||
layer_id,
|
||||
)
|
||||
for layer_id in range(config.n_layers)
|
||||
]
|
||||
[DecoderBlock(config, layer_id) for layer_id in range(config.n_layers)]
|
||||
)
|
||||
|
||||
self.norm = RMSNorm(config.dim, config.norm_eps)
|
||||
|
||||
@@ -5,7 +5,7 @@ import torch.nn as nn
|
||||
from torch import Tensor
|
||||
|
||||
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.components.decoder_block import DecoderBlock
|
||||
from astrai.model.components.embedding import Embedding
|
||||
@@ -59,31 +59,12 @@ class AutoRegressiveLM(AutoModel):
|
||||
self.rotary_embedding = RotaryEmbedding(
|
||||
rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling
|
||||
)
|
||||
self.embed_tokens = Embedding(config.vocab_size, config.dim)
|
||||
self.embed_tokens = Embedding(
|
||||
config.vocab_size, config.dim, neftune_alpha=config.neftune_alpha
|
||||
)
|
||||
|
||||
self.layers = nn.ModuleList(
|
||||
[
|
||||
DecoderBlock(
|
||||
config.dim,
|
||||
config.n_heads,
|
||||
config.dim_ffn,
|
||||
config.n_kv_heads,
|
||||
config.norm_eps,
|
||||
config.use_qk_norm,
|
||||
config.use_gated_attention,
|
||||
layer_id,
|
||||
attn_type=config.attn_type,
|
||||
ffn_type=config.ffn_type,
|
||||
n_routed_experts=config.n_routed_experts,
|
||||
n_shared_experts=config.n_shared_experts,
|
||||
n_activated_experts=config.n_activated_experts,
|
||||
topk_method=config.topk_method,
|
||||
kv_lora_rank=config.kv_lora_rank,
|
||||
qk_nope_head_dim=config.qk_nope_head_dim,
|
||||
qk_rope_head_dim=config.qk_rope_head_dim,
|
||||
)
|
||||
for layer_id in range(config.n_layers)
|
||||
]
|
||||
[DecoderBlock(config, layer_id) for layer_id in range(config.n_layers)]
|
||||
)
|
||||
|
||||
self.norm = RMSNorm(config.dim, config.norm_eps)
|
||||
@@ -131,7 +112,7 @@ class AutoRegressiveLM(AutoModel):
|
||||
self,
|
||||
input_ids: Tensor,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
paged_cache: Optional[KvcacheView] = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
position_ids: Optional[Tensor] = None,
|
||||
) -> Dict[str, Tensor]:
|
||||
assert input_ids.ndim == 2
|
||||
|
||||
@@ -7,6 +7,7 @@ from contextlib import contextmanager
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
from torch.distributed.fsdp import FullStateDictConfig, StateDictType
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
@@ -120,6 +121,21 @@ class BaseExecutor:
|
||||
def unwrap_model(self, model: nn.Module):
|
||||
return model.state_dict()
|
||||
|
||||
@contextmanager
|
||||
def checkpoint_context(self, model: nn.Module):
|
||||
if self.use_distributed:
|
||||
dist.barrier()
|
||||
state_dict = self._gather_state_dict(model)
|
||||
yield state_dict
|
||||
if self.use_distributed:
|
||||
dist.barrier()
|
||||
|
||||
def _gather_state_dict(self, model: nn.Module):
|
||||
state_dict = self.unwrap_model(model)
|
||||
if self.use_distributed and get_rank() != 0:
|
||||
return None
|
||||
return state_dict
|
||||
|
||||
@property
|
||||
def use_distributed(self) -> bool:
|
||||
return get_world_size() > 1
|
||||
@@ -132,6 +148,19 @@ class BaseExecutor:
|
||||
def grad_accum_steps(self) -> int:
|
||||
return self.gradient_state.num_steps
|
||||
|
||||
def clip_grad_norm(self, model: nn.Module, max_norm: Optional[float]) -> float:
|
||||
if max_norm is None:
|
||||
total_norm = torch.norm(
|
||||
torch.stack(
|
||||
[p.grad.norm(2) for p in model.parameters() if p.grad is not None]
|
||||
)
|
||||
)
|
||||
return total_norm.item()
|
||||
total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)
|
||||
if isinstance(total_norm, torch.Tensor):
|
||||
return total_norm.item()
|
||||
return total_norm
|
||||
|
||||
|
||||
class ExecutorFactory(BaseFactory[BaseExecutor]):
|
||||
pass
|
||||
@@ -260,12 +289,22 @@ class FSDPExecutor(BaseExecutor):
|
||||
return model.no_sync()
|
||||
return contextlib.nullcontext()
|
||||
|
||||
def clip_grad_norm(self, model: nn.Module, max_norm: Optional[float]) -> float:
|
||||
if max_norm is None:
|
||||
return super().clip_grad_norm(model, max_norm)
|
||||
if isinstance(model, FSDP) and self.use_distributed:
|
||||
total_norm = model.clip_grad_norm_(max_norm)
|
||||
if isinstance(total_norm, torch.Tensor):
|
||||
return total_norm.item()
|
||||
return total_norm
|
||||
return super().clip_grad_norm(model, max_norm)
|
||||
|
||||
def unwrap_model(self, model: nn.Module):
|
||||
if isinstance(model, FSDP) and self.use_distributed:
|
||||
with FSDP.state_dict_type(
|
||||
model,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=False),
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
):
|
||||
return model.state_dict()
|
||||
|
||||
|
||||
@@ -1,14 +1,21 @@
|
||||
import os
|
||||
import socket
|
||||
from abc import ABC, abstractmethod
|
||||
from contextlib import contextmanager
|
||||
from functools import wraps
|
||||
from typing import Callable
|
||||
from typing import Callable, Optional
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
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():
|
||||
return os.environ["LOCAL_DEVICE"]
|
||||
|
||||
@@ -58,9 +65,11 @@ def setup_parallel(
|
||||
os.environ["WORLD_SIZE"] = str(world_size)
|
||||
os.environ["LOCAL_DEVICE"] = str(device_id)
|
||||
|
||||
dist.init_process_group(
|
||||
rank=rank, world_size=world_size, backend=backend, device_id=device_id
|
||||
)
|
||||
pg_kwargs = dict(rank=rank, world_size=world_size, backend=backend)
|
||||
if backend in ("nccl", "ccl"):
|
||||
pg_kwargs["device_id"] = device_id
|
||||
|
||||
dist.init_process_group(**pg_kwargs)
|
||||
|
||||
try:
|
||||
if backend == "nccl" and torch.cuda.is_available():
|
||||
@@ -215,11 +224,13 @@ def spawn_parallel_fn(
|
||||
world_size: int,
|
||||
backend: str = "nccl",
|
||||
master_addr: str = "localhost",
|
||||
master_port: str = "29500",
|
||||
master_port: Optional[str] = None,
|
||||
device_type: str = "cuda",
|
||||
start_method: str = "spawn",
|
||||
**kwargs,
|
||||
):
|
||||
if master_port is None:
|
||||
master_port = find_free_port()
|
||||
launcher = _detect_launcher()
|
||||
if launcher in ("torchelastic", "torchrun", "external"):
|
||||
strategy = TorchrunStrategy(
|
||||
|
||||
@@ -1,17 +1,21 @@
|
||||
from astrai.preprocessing.builder import (
|
||||
BaseMaskBuilder,
|
||||
MaskBuilderFactory,
|
||||
MultiOutputMaskBuilder,
|
||||
SectionedMaskBuilder,
|
||||
SingleOutputMaskBuilder,
|
||||
)
|
||||
from astrai.preprocessing.packing import (
|
||||
PackingStrategy,
|
||||
PackingStrategyFactory,
|
||||
plan_bfd,
|
||||
)
|
||||
from astrai.preprocessing.pipeline import Pipeline, filter_by_length
|
||||
from astrai.preprocessing.position_id import (
|
||||
PositionIdStrategy,
|
||||
PositionIdStrategyFactory,
|
||||
)
|
||||
from astrai.preprocessing.transform import TokenizeTransform
|
||||
from astrai.preprocessing.writer import (
|
||||
StoreWriter,
|
||||
StoreWriterFactory,
|
||||
@@ -20,13 +24,17 @@ from astrai.preprocessing.writer import (
|
||||
__all__ = [
|
||||
"BaseMaskBuilder",
|
||||
"MaskBuilderFactory",
|
||||
"MultiOutputMaskBuilder",
|
||||
"PackingStrategy",
|
||||
"PackingStrategyFactory",
|
||||
"Pipeline",
|
||||
"PositionIdStrategy",
|
||||
"PositionIdStrategyFactory",
|
||||
"SectionedMaskBuilder",
|
||||
"SingleOutputMaskBuilder",
|
||||
"StoreWriter",
|
||||
"StoreWriterFactory",
|
||||
"TokenizeTransform",
|
||||
"filter_by_length",
|
||||
"plan_bfd",
|
||||
]
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
"""Mask building for preprocessing pipeline.
|
||||
|
||||
:class:`SectionRenderer` converts section specs into token ids and loss
|
||||
masks (template / text / value extraction). :class:`SectionedMaskBuilder`
|
||||
orchestrates single-output / multi-output (DPO / GRPO) assembly.
|
||||
masks (template / text / value extraction). :class:`SingleOutputMaskBuilder`
|
||||
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
|
||||
@@ -93,8 +95,15 @@ class SectionRenderer:
|
||||
return all_ids, loss_mask
|
||||
|
||||
def process_list_field(self, item: dict, sections: list, config, tokenizer):
|
||||
all_ids: list[int] = []
|
||||
loss_mask: list[int] = []
|
||||
"""Tokenize a list-valued field, preserving per-element boundaries.
|
||||
|
||||
Returns ``(list_of_id_lists, list_of_mask_lists)`` where each
|
||||
inner list corresponds to one element of the source list. This
|
||||
is critical for GRPO where each response must stay a separate
|
||||
sequence so the strategy can form a ``[G, R]`` tensor.
|
||||
"""
|
||||
per_item_ids: list[list[int]] = []
|
||||
per_item_masks: list[list[int]] = []
|
||||
|
||||
for sec in sections:
|
||||
field = sec["field"]
|
||||
@@ -106,17 +115,13 @@ class SectionRenderer:
|
||||
continue
|
||||
|
||||
for val in values:
|
||||
ids: list[int] = []
|
||||
mask: list[int] = []
|
||||
if use_template:
|
||||
if isinstance(val, list):
|
||||
wrapper = {field: val}
|
||||
self._append_template(
|
||||
wrapper,
|
||||
field,
|
||||
action,
|
||||
tokenizer,
|
||||
config,
|
||||
all_ids,
|
||||
loss_mask,
|
||||
wrapper, field, action, tokenizer, config, ids, mask
|
||||
)
|
||||
else:
|
||||
wrapper = {field: str(val)}
|
||||
@@ -128,17 +133,19 @@ class SectionRenderer:
|
||||
False,
|
||||
False,
|
||||
config,
|
||||
all_ids,
|
||||
loss_mask,
|
||||
ids,
|
||||
mask,
|
||||
)
|
||||
if ids:
|
||||
max_len = config.preprocessing.max_seq_len
|
||||
ids = ids[:max_len]
|
||||
mask = mask[: len(ids)]
|
||||
per_item_ids.append(ids)
|
||||
per_item_masks.append(mask)
|
||||
|
||||
max_len = config.preprocessing.max_seq_len
|
||||
all_ids = all_ids[:max_len]
|
||||
loss_mask = loss_mask[: len(all_ids)]
|
||||
|
||||
if not all_ids:
|
||||
if not per_item_ids:
|
||||
return None, None
|
||||
return all_ids, loss_mask
|
||||
return per_item_ids, per_item_masks
|
||||
|
||||
@staticmethod
|
||||
def is_value_section(sections: list) -> bool:
|
||||
@@ -212,42 +219,17 @@ class MaskBuilderFactory(BaseFactory["BaseMaskBuilder"]):
|
||||
pass
|
||||
|
||||
|
||||
@MaskBuilderFactory.register("sectioned")
|
||||
class SectionedMaskBuilder(BaseMaskBuilder):
|
||||
"""Config-driven builder supporting single and multi-output modes.
|
||||
@MaskBuilderFactory.register("single")
|
||||
class SingleOutputMaskBuilder(BaseMaskBuilder):
|
||||
"""Build a single output sequence with optional loss mask.
|
||||
|
||||
Single-output::
|
||||
|
||||
{"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"``)
|
||||
Expects ``config.input.sections`` (list of section specs).
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.renderer = SectionRenderer()
|
||||
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 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
|
||||
if not sections:
|
||||
return None
|
||||
@@ -266,9 +248,22 @@ class SectionedMaskBuilder(BaseMaskBuilder):
|
||||
result["loss_mask"] = mask
|
||||
return result
|
||||
|
||||
def _build_multi(
|
||||
self, item: dict, sources_spec: dict, config, tokenizer
|
||||
) -> Optional[dict]:
|
||||
|
||||
@MaskBuilderFactory.register("multi")
|
||||
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 = {}
|
||||
any_output = False
|
||||
|
||||
@@ -292,10 +287,18 @@ class SectionedMaskBuilder(BaseMaskBuilder):
|
||||
ids, mask = self.renderer.process_list_field(
|
||||
item, sections, config, tokenizer
|
||||
)
|
||||
else:
|
||||
ids, mask = self.renderer.process_sections(
|
||||
item, sections, config, tokenizer, is_top_level=True
|
||||
)
|
||||
if ids is None:
|
||||
continue
|
||||
# ids is List[List[int]] — preserve per-response structure
|
||||
result[output_key] = ids
|
||||
if mask is not None:
|
||||
result[mask_key] = mask
|
||||
any_output = True
|
||||
continue
|
||||
|
||||
ids, mask = self.renderer.process_sections(
|
||||
item, sections, config, tokenizer, is_top_level=True
|
||||
)
|
||||
|
||||
if ids is None:
|
||||
continue
|
||||
@@ -313,3 +316,22 @@ class SectionedMaskBuilder(BaseMaskBuilder):
|
||||
|
||||
result["domain"] = _extract_domain(item, config.output.domain_key)
|
||||
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)
|
||||
|
||||
@@ -0,0 +1,124 @@
|
||||
"""Shared preprocessing kernel used by both :class:`Pipeline` and
|
||||
:class:`TokenizeTransform`.
|
||||
|
||||
The two entry points previously duplicated ~60 % of their logic:
|
||||
record iteration, mask-builder invocation, primary-id extraction,
|
||||
per-key accumulation, dtype inference and position-id generation.
|
||||
This module factors out the common core as pure functions so that
|
||||
the online (``TokenizeTransform``) and offline (``Pipeline``) paths
|
||||
stay in lockstep.
|
||||
"""
|
||||
|
||||
from itertools import chain
|
||||
from typing import Dict, Iterator, List, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.config.preprocess_config import PipelineConfig
|
||||
from astrai.preprocessing.builder import MaskBuilderFactory
|
||||
from astrai.preprocessing.position_id import PositionIdStrategyFactory
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
|
||||
def build_preprocessing_components(config: PipelineConfig, tokenizer_path: str):
|
||||
"""Load tokenizer, mask builder and position-id strategy together.
|
||||
|
||||
Both ``Pipeline`` and ``TokenizeTransform`` need the same triple;
|
||||
centralising the construction avoids drift (e.g. one path forgetting
|
||||
to create the position-id strategy).
|
||||
"""
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
|
||||
mask_builder = MaskBuilderFactory.create("sectioned")
|
||||
position_strategy = PositionIdStrategyFactory.create(
|
||||
config.output.position_ids_mode
|
||||
)
|
||||
return tokenizer, mask_builder, position_strategy
|
||||
|
||||
|
||||
def primary_ids(result: dict) -> List[int]:
|
||||
"""Return the first flat int-list value in *result*.
|
||||
|
||||
Used for token counting and position-id generation when the
|
||||
primary key name is not known (DPO uses ``chosen``, GRPO uses
|
||||
``prompts``, SFT uses ``sequence``).
|
||||
"""
|
||||
for val in result.values():
|
||||
if isinstance(val, list) and val and isinstance(val[0], int):
|
||||
return val
|
||||
return []
|
||||
|
||||
|
||||
def infer_dtype(ids: List) -> torch.dtype:
|
||||
"""Float values become float32, everything else int32."""
|
||||
if ids and isinstance(ids[0], float):
|
||||
return torch.float32
|
||||
return torch.int32
|
||||
|
||||
|
||||
def iter_raw_records(
|
||||
records: List[dict],
|
||||
mask_builder,
|
||||
config: PipelineConfig,
|
||||
tokenizer,
|
||||
) -> Iterator[dict]:
|
||||
"""Yield mask-builder output dicts for each record, skipping failures.
|
||||
|
||||
Drops ``domain`` from the result (callers that need it should read
|
||||
it before calling this). Each yielded dict maps a key
|
||||
(``sequence``, ``chosen``, ``responses``…) to either a flat
|
||||
``List[int]`` or a nested ``List[List[int]]`` (GRPO responses/masks).
|
||||
"""
|
||||
for item in records:
|
||||
result = mask_builder.build(item, config, tokenizer)
|
||||
if result is None:
|
||||
continue
|
||||
result.pop("domain", None)
|
||||
if not primary_ids(result):
|
||||
continue
|
||||
yield result
|
||||
|
||||
|
||||
def to_per_record_tensors(
|
||||
raw: Dict[str, list],
|
||||
) -> Dict[str, List[torch.Tensor]]:
|
||||
"""Convert an accumulated ``{key: [per-record ids]}`` dict to tensors.
|
||||
|
||||
Handles three shapes transparently:
|
||||
|
||||
- ``List[int]`` per record (``sequence``, ``chosen``…) → one tensor per record.
|
||||
- ``List[List[int]]`` per record (GRPO ``responses``/``masks``) → one
|
||||
``List[Tensor]`` per record (nested), preserving the per-response
|
||||
boundary so downstream code can index responses individually.
|
||||
- ``List[int]`` for the whole shard (pre-packed keys) → single tensor.
|
||||
|
||||
The detection mirrors the previous inline logic in
|
||||
``Pipeline._flush`` and ``TokenizeTransform.apply``.
|
||||
"""
|
||||
tensors: Dict[str, List[torch.Tensor]] = {}
|
||||
for key, ids_list in raw.items():
|
||||
if ids_list and isinstance(ids_list[0], list):
|
||||
tensors[key] = [
|
||||
[torch.tensor(sub, dtype=infer_dtype(sub)) for sub in ids]
|
||||
if ids and isinstance(ids[0], list)
|
||||
else torch.tensor(ids, dtype=infer_dtype(ids))
|
||||
for ids in ids_list
|
||||
]
|
||||
else:
|
||||
tensors[key] = [
|
||||
torch.tensor(list(chain.from_iterable(ids_list)), dtype=torch.int32)
|
||||
]
|
||||
return tensors
|
||||
|
||||
|
||||
def build_position_ids(
|
||||
sequences: List[List[int]],
|
||||
strategy,
|
||||
) -> Optional[List[int]]:
|
||||
"""Generate position ids for *sequences* using *strategy*.
|
||||
|
||||
Returns ``None`` when the strategy produces no ids (e.g. ``none``
|
||||
mode), so callers can skip attaching the key instead of storing
|
||||
an empty list.
|
||||
"""
|
||||
pos_ids = strategy.generate(sequences)
|
||||
return pos_ids or None
|
||||
@@ -19,6 +19,43 @@ def _truncate(seq: List[int], max_len: int, mode: str) -> List[int]:
|
||||
return seq[:max_len]
|
||||
|
||||
|
||||
def plan_bfd(
|
||||
sequences: List[List[int]], max_packed_len: int, truncation_mode: str = "keep_start"
|
||||
) -> List[List[int]]:
|
||||
"""Best-Fit Decreasing bin packing of *sequences* into bins.
|
||||
|
||||
Returns a list of bins, each bin a list of original indices into
|
||||
*sequences*. Bin capacities are respected on the *truncated*
|
||||
length of each sequence (so a sequence longer than
|
||||
*max_packed_len* counts at *max_packed_len*).
|
||||
|
||||
Pure index-based so callers can apply the same plan to any
|
||||
aligned key (``loss_mask``, ``position_ids``…).
|
||||
"""
|
||||
n = len(sequences)
|
||||
order = sorted(range(n), key=lambda i: len(sequences[i]), reverse=True)
|
||||
bins: List[List[int]] = []
|
||||
bin_lengths: List[int] = []
|
||||
|
||||
for orig_idx in order:
|
||||
seq_len = len(_truncate(sequences[orig_idx], max_packed_len, truncation_mode))
|
||||
best_bin = None
|
||||
best_remain = max_packed_len + 1
|
||||
for i, bl in enumerate(bin_lengths):
|
||||
remain = max_packed_len - bl
|
||||
if seq_len <= remain < best_remain:
|
||||
best_remain = remain
|
||||
best_bin = i
|
||||
if best_bin is not None:
|
||||
bins[best_bin].append(orig_idx)
|
||||
bin_lengths[best_bin] += seq_len
|
||||
else:
|
||||
bins.append([orig_idx])
|
||||
bin_lengths.append(seq_len)
|
||||
|
||||
return bins
|
||||
|
||||
|
||||
class PackingStrategy(ABC):
|
||||
"""Reorder and truncate sequences within a shard."""
|
||||
|
||||
@@ -70,7 +107,7 @@ class BFDPacking(PackingStrategy):
|
||||
sequences = keys.get("sequence", [])
|
||||
if not sequences:
|
||||
return keys
|
||||
bins = self._plan(sequences, max_packed_len, truncation_mode)
|
||||
bins = plan_bfd(sequences, max_packed_len, truncation_mode)
|
||||
|
||||
packed: Dict[str, List[List[int]]] = {}
|
||||
for k, vals in keys.items():
|
||||
@@ -91,31 +128,49 @@ class BFDPacking(PackingStrategy):
|
||||
result.extend(vals[i])
|
||||
return result
|
||||
|
||||
|
||||
@PackingStrategyFactory.register("bfd_split")
|
||||
class BFDSplitPacking(BFDPacking):
|
||||
"""BFD packing with over-length sequences split into chunks.
|
||||
|
||||
Sequences longer than *max_packed_len* are split into consecutive
|
||||
chunks of at most *max_packed_len* tokens instead of being
|
||||
truncated. Each chunk becomes an independent sequence that enters
|
||||
BFD planning. All keys (``loss_mask``, ``position_ids``, …) are
|
||||
split in lockstep so per-token alignment is preserved.
|
||||
|
||||
Note: because each chunk is treated as a separate document, the
|
||||
second chunk of a split sequence loses the preceding context.
|
||||
"""
|
||||
|
||||
def apply(
|
||||
self,
|
||||
keys: Dict[str, List[List[int]]],
|
||||
max_packed_len: int,
|
||||
truncation_mode: str,
|
||||
) -> Dict[str, List[List[int]]]:
|
||||
sequences = keys.get("sequence", [])
|
||||
if not sequences:
|
||||
return keys
|
||||
if max_packed_len <= 0:
|
||||
return super().apply(keys, max_packed_len, truncation_mode)
|
||||
|
||||
split_keys = self._split_all(keys, max_packed_len)
|
||||
return super().apply(split_keys, max_packed_len, truncation_mode)
|
||||
|
||||
@staticmethod
|
||||
def _plan(
|
||||
sequences: List[List[int]], max_packed_len: int, truncation_mode: str
|
||||
) -> List[List[int]]:
|
||||
n = len(sequences)
|
||||
order = sorted(range(n), key=lambda i: len(sequences[i]), reverse=True)
|
||||
bins: List[List[int]] = []
|
||||
bin_lengths: List[int] = []
|
||||
|
||||
for orig_idx in order:
|
||||
seq_len = len(
|
||||
_truncate(sequences[orig_idx], max_packed_len, truncation_mode)
|
||||
)
|
||||
best_bin = None
|
||||
best_remain = max_packed_len + 1
|
||||
for i, bl in enumerate(bin_lengths):
|
||||
remain = max_packed_len - bl
|
||||
if seq_len <= remain < best_remain:
|
||||
best_remain = remain
|
||||
best_bin = i
|
||||
if best_bin is not None:
|
||||
bins[best_bin].append(orig_idx)
|
||||
bin_lengths[best_bin] += seq_len
|
||||
else:
|
||||
bins.append([orig_idx])
|
||||
bin_lengths.append(seq_len)
|
||||
|
||||
return bins
|
||||
def _split_all(
|
||||
keys: Dict[str, List[List[int]]], max_packed_len: int
|
||||
) -> Dict[str, List[List[int]]]:
|
||||
"""Split every sequence exceeding *max_packed_len* into chunks,
|
||||
applying the same chunk boundaries to all keys."""
|
||||
sequences = keys["sequence"]
|
||||
chunk_bounds = [list(range(0, len(s), max_packed_len)) for s in sequences]
|
||||
result: Dict[str, List[List[int]]] = {}
|
||||
for key, vals in keys.items():
|
||||
split_vals: List[List[int]] = []
|
||||
for val, starts in zip(vals, chunk_bounds):
|
||||
for start in starts:
|
||||
split_vals.append(val[start : start + max_packed_len])
|
||||
result[key] = split_vals
|
||||
return result
|
||||
|
||||
@@ -4,6 +4,10 @@ Composes a :class:`BaseMaskBuilder` (selected by ``input.type``) with
|
||||
sharding and flush to ``.h5`` / ``.bin`` storage. Packing, position-id
|
||||
generation and storage writing are each delegated to pluggable strategies,
|
||||
dispatched by configuration keys.
|
||||
|
||||
Record iteration, mask building, primary-id extraction and per-key
|
||||
accumulation are shared with :class:`TokenizeTransform` via the
|
||||
:mod:`astrai.preprocessing.core` helpers.
|
||||
"""
|
||||
|
||||
import json
|
||||
@@ -17,11 +21,13 @@ import torch
|
||||
import tqdm
|
||||
|
||||
from astrai.config.preprocess_config import PipelineConfig
|
||||
from astrai.preprocessing.builder import MaskBuilderFactory
|
||||
from astrai.preprocessing.core import (
|
||||
build_preprocessing_components,
|
||||
iter_raw_records,
|
||||
primary_ids,
|
||||
)
|
||||
from astrai.preprocessing.packing import PackingStrategyFactory
|
||||
from astrai.preprocessing.position_id import PositionIdStrategyFactory
|
||||
from astrai.preprocessing.writer import StoreWriterFactory
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -64,20 +70,18 @@ class Pipeline:
|
||||
self.output_dir = output_dir
|
||||
self.tokenizer_path = tokenizer_path
|
||||
|
||||
self.mask_builder = MaskBuilderFactory.create("sectioned")
|
||||
self.tokenizer, self.mask_builder, self._position_id = (
|
||||
build_preprocessing_components(config, tokenizer_path)
|
||||
)
|
||||
self._packer = PackingStrategyFactory.create(
|
||||
config.preprocessing.packing_strategy
|
||||
)
|
||||
self._position_id = PositionIdStrategyFactory.create(
|
||||
config.output.position_ids_mode
|
||||
)
|
||||
self._writer = StoreWriterFactory.create(config.output.storage_format)
|
||||
|
||||
def transform(self, item: dict) -> Optional[dict]:
|
||||
return self.mask_builder.build(item, self.config, self._tokenizer)
|
||||
return self.mask_builder.build(item, self.config, self.tokenizer)
|
||||
|
||||
def run(self):
|
||||
self._tokenizer = AutoTokenizer.from_pretrained(self.tokenizer_path)
|
||||
domains: dict = defaultdict(lambda: defaultdict(list))
|
||||
total_tokens = 0
|
||||
shard_idx: dict[str, int] = defaultdict(int)
|
||||
@@ -102,14 +106,7 @@ class Pipeline:
|
||||
continue
|
||||
|
||||
domain = result.pop("domain", "__default__")
|
||||
|
||||
is_multi = bool(getattr(self.config.input, "sources", None))
|
||||
if is_multi:
|
||||
ids = self._primary_ids(result)
|
||||
else:
|
||||
ids = result.pop("sequence")
|
||||
result["sequence"] = ids
|
||||
|
||||
ids = primary_ids(result)
|
||||
if not ids:
|
||||
continue
|
||||
|
||||
@@ -129,51 +126,44 @@ class Pipeline:
|
||||
if total_tokens > 0:
|
||||
self._flush(domains, shard_idx)
|
||||
|
||||
@staticmethod
|
||||
def _primary_ids(result: dict) -> list:
|
||||
"""Return the first list-valued entry in *result* as the primary id
|
||||
sequence for token counting."""
|
||||
for val in result.values():
|
||||
if isinstance(val, list) and val and isinstance(val[0], int):
|
||||
return val
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
def _align_bucket(bucket: dict, result: dict, ids: list):
|
||||
"""Pad previously-accumulated keys that are missing from *result*."""
|
||||
for key in list(bucket.keys()):
|
||||
if key in result:
|
||||
continue
|
||||
bucket[key].append([1] * len(ids))
|
||||
bucket[key].append([0] * len(ids))
|
||||
|
||||
def _iter_items(self):
|
||||
for path in self.paths:
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
yield json.loads(line)
|
||||
if path.endswith(".json"):
|
||||
data = json.load(f)
|
||||
if isinstance(data, dict):
|
||||
yield data
|
||||
elif isinstance(data, list):
|
||||
yield from data
|
||||
else:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
yield json.loads(line)
|
||||
|
||||
def _flush(self, domains, shard_idx):
|
||||
for domain, keys in domains.items():
|
||||
idx = shard_idx[domain]
|
||||
|
||||
pp = self.config.preprocessing
|
||||
original_sequences = keys.get("sequence", [])
|
||||
mode = self.config.output.position_ids_mode
|
||||
|
||||
keys = self._inject_doc_reset_position_ids(keys, mode, original_sequences)
|
||||
keys = self._packer.apply(dict(keys), pp.max_packed_len, pp.truncation_mode)
|
||||
|
||||
tensors: Dict[str, List[torch.Tensor]] = {}
|
||||
for key, ids_list in keys.items():
|
||||
dt = _STR_TO_DTYPE.get(
|
||||
self.config.output.dtype.get(key, "int32"), torch.int32
|
||||
)
|
||||
tensors[key] = [
|
||||
torch.tensor(list(chain.from_iterable(ids_list)), dtype=dt)
|
||||
]
|
||||
|
||||
pos_ids = self._position_id.generate(keys.get("sequence", []))
|
||||
if pos_ids:
|
||||
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
|
||||
tensors = self._to_tensors(keys)
|
||||
tensors = self._inject_continuous_position_ids(
|
||||
tensors, mode, keys.get("sequence", [])
|
||||
)
|
||||
|
||||
self._writer.save(self.output_dir, domain, idx, tensors)
|
||||
shard_idx[domain] = idx + 1
|
||||
@@ -183,3 +173,76 @@ class Pipeline:
|
||||
f" saved {domain}/shard_{idx:04d} "
|
||||
f"({tensors[first_key][0].numel():,} tokens)"
|
||||
)
|
||||
|
||||
def _inject_doc_reset_position_ids(
|
||||
self,
|
||||
keys: Dict[str, list],
|
||||
mode: str,
|
||||
original_sequences: List[List[int]],
|
||||
) -> Dict[str, list]:
|
||||
"""Attach per-document position_ids before packing (``doc_reset``).
|
||||
|
||||
``doc_reset`` position ids must enter the packer so that each
|
||||
packed bin concatenates the per-doc ranges in bin order. The
|
||||
per-record structure ``[range(len(s)) for s in seqs]`` is required
|
||||
by the packer (it concatenates per-record lists per bin); the
|
||||
``PositionIdStrategy.generate`` flattens, so it cannot be used
|
||||
directly here — it is only consulted for the ``continuous``
|
||||
post-packing path.
|
||||
"""
|
||||
if mode != "doc_reset" or not original_sequences:
|
||||
return keys
|
||||
keys["position_ids"] = [list(range(len(s))) for s in original_sequences]
|
||||
return keys
|
||||
|
||||
def _inject_continuous_position_ids(
|
||||
self,
|
||||
tensors: Dict[str, List[torch.Tensor]],
|
||||
mode: str,
|
||||
packed_sequences: List[List[int]],
|
||||
) -> Dict[str, List[torch.Tensor]]:
|
||||
"""Attach a single continuous position_ids tensor after packing.
|
||||
|
||||
``continuous`` mode spans the whole shard (post-packing), so it
|
||||
cannot participate in bin packing — it is computed from the
|
||||
packed sequences and appended directly to the tensor dict.
|
||||
"""
|
||||
if mode != "continuous" or not packed_sequences:
|
||||
return tensors
|
||||
pos_ids = self._position_id.generate(packed_sequences)
|
||||
if pos_ids:
|
||||
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
|
||||
return tensors
|
||||
|
||||
def _to_tensors(self, keys: Dict[str, list]) -> Dict[str, List[torch.Tensor]]:
|
||||
"""Convert packed per-key id lists to tensors.
|
||||
|
||||
Honours ``config.output.dtype`` overrides per key; falls back to
|
||||
``int32``. Handles three shapes (see
|
||||
:func:`astrai.preprocessing.core.to_per_record_tensors` for the
|
||||
equivalent online-path helper):
|
||||
- ``List[int]`` per record → one tensor per record.
|
||||
- ``List[List[int]]`` per record (GRPO responses/masks) → one tensor
|
||||
per record, inner lists flattened.
|
||||
- ``List[int]`` for the whole shard (pre-packed keys) → single tensor.
|
||||
"""
|
||||
tensors: Dict[str, List[torch.Tensor]] = {}
|
||||
for key, ids_list in keys.items():
|
||||
dt = _STR_TO_DTYPE.get(
|
||||
self.config.output.dtype.get(key, "int32"), torch.int32
|
||||
)
|
||||
if ids_list and isinstance(ids_list[0], list):
|
||||
tensors[key] = [
|
||||
torch.tensor(
|
||||
list(chain.from_iterable(ids))
|
||||
if ids and isinstance(ids[0], list)
|
||||
else ids,
|
||||
dtype=dt,
|
||||
)
|
||||
for ids in ids_list
|
||||
]
|
||||
else:
|
||||
tensors[key] = [
|
||||
torch.tensor(list(chain.from_iterable(ids_list)), dtype=dt)
|
||||
]
|
||||
return tensors
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
"""Tokenization transform for JSONL record streams.
|
||||
|
||||
Bridges the Reader layer (``JsonlStore`` reads raw JSON records) and the
|
||||
Dataset layer (expects per-record tensors). Holds the tokenizer,
|
||||
mask-builder and position-id strategy together so that I/O code stays
|
||||
free of model dependencies.
|
||||
|
||||
The record-processing core (mask building, primary-id extraction,
|
||||
per-key tensorisation, position-id generation) is shared with
|
||||
:class:`astrai.preprocessing.pipeline.Pipeline` via the
|
||||
:mod:`astrai.preprocessing.core` helpers.
|
||||
"""
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Dict, List
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.config.preprocess_config import PipelineConfig
|
||||
from astrai.preprocessing.core import (
|
||||
build_position_ids,
|
||||
build_preprocessing_components,
|
||||
iter_raw_records,
|
||||
to_per_record_tensors,
|
||||
)
|
||||
|
||||
|
||||
class TokenizeTransform:
|
||||
"""Tokenize raw JSONL record dicts into per-key tensor lists.
|
||||
|
||||
Owns the three preprocessing concerns that were previously inlined in
|
||||
``JsonlStore``: tokenization, loss-mask construction and position-id
|
||||
generation. Constructing it loads the tokenizer, so it is intentionally
|
||||
cheap to pass around once built.
|
||||
|
||||
Args:
|
||||
config: Pipeline config describing sections / masks / position mode.
|
||||
tokenizer_path: Path passed to ``AutoTokenizer.from_pretrained``.
|
||||
"""
|
||||
|
||||
def __init__(self, config: PipelineConfig, tokenizer_path: str):
|
||||
self.config = config
|
||||
self.tokenizer, self.mask_builder, self.position_strategy = (
|
||||
build_preprocessing_components(config, tokenizer_path)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_config_file(cls, config_path: str) -> "TokenizeTransform":
|
||||
"""Build from a ``dataset_config.json`` file path.
|
||||
|
||||
The config file follows :class:`PipelineConfig` schema with an
|
||||
extra ``tokenizer_path`` field. When omitted, the config's
|
||||
parent directory is used as the tokenizer path.
|
||||
"""
|
||||
root = Path(config_path).parent
|
||||
with open(config_path, "r", encoding="utf-8") as f:
|
||||
raw_config = json.load(f)
|
||||
tokenizer_path = raw_config.pop("tokenizer_path", None) or str(root)
|
||||
config = PipelineConfig.from_dict(raw_config)
|
||||
return cls(config, tokenizer_path)
|
||||
|
||||
def apply(self, records: List[dict]) -> Dict[str, list]:
|
||||
"""Tokenize a list of raw record dicts.
|
||||
|
||||
Returns a dict mapping key (``sequence``, ``chosen``, ``responses``,
|
||||
…) to a list of per-record tensors (or nested tensor lists for
|
||||
multi-response keys such as GRPO ``responses``).
|
||||
"""
|
||||
raw: Dict[str, list] = {}
|
||||
doc_sequences: List[List[int]] = []
|
||||
|
||||
for result in iter_raw_records(
|
||||
records, self.mask_builder, self.config, self.tokenizer
|
||||
):
|
||||
primary = None
|
||||
for val in result.values():
|
||||
if isinstance(val, list) and val and isinstance(val[0], int):
|
||||
primary = val
|
||||
break
|
||||
if primary is not None:
|
||||
doc_sequences.append(primary)
|
||||
for key, ids in result.items():
|
||||
raw.setdefault(key, []).append(ids)
|
||||
|
||||
tensors = to_per_record_tensors(raw)
|
||||
|
||||
pos_ids = build_position_ids(doc_sequences, self.position_strategy)
|
||||
if pos_ids is not None:
|
||||
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
|
||||
|
||||
return tensors
|
||||
@@ -14,8 +14,8 @@ from typing import Dict, List
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.dataset.storage import save_bin, save_h5
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.serialization import save_bin, save_h5
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
"""Serialization utilities for models and datasets.
|
||||
|
||||
This package re-exports checkpoint helpers and dataset storage helpers so
|
||||
that existing imports from ``astrai.serialization`` continue to work.
|
||||
"""
|
||||
|
||||
from astrai.serialization.checkpoint import (
|
||||
Checkpoint,
|
||||
load_json,
|
||||
load_model_config,
|
||||
load_model_weights,
|
||||
load_safetensors,
|
||||
load_state_dict,
|
||||
load_torch,
|
||||
save_json,
|
||||
save_model,
|
||||
save_safetensors,
|
||||
save_torch,
|
||||
)
|
||||
from astrai.serialization.dataset import (
|
||||
load_bin,
|
||||
load_bin_offsets,
|
||||
load_h5,
|
||||
save_bin,
|
||||
save_h5,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"Checkpoint",
|
||||
"load_json",
|
||||
"load_model_config",
|
||||
"load_model_weights",
|
||||
"load_safetensors",
|
||||
"load_state_dict",
|
||||
"load_torch",
|
||||
"save_json",
|
||||
"save_model",
|
||||
"save_safetensors",
|
||||
"save_torch",
|
||||
"load_bin",
|
||||
"load_bin_offsets",
|
||||
"load_h5",
|
||||
"save_bin",
|
||||
"save_h5",
|
||||
]
|
||||
@@ -1,3 +1,5 @@
|
||||
"""Model checkpoint serialization helpers."""
|
||||
|
||||
import io
|
||||
import json
|
||||
import time
|
||||
@@ -136,7 +138,7 @@ def load_state_dict(path: Union[str, Path], broadcast: bool = False) -> dict:
|
||||
class Checkpoint:
|
||||
state_dict: Dict[str, Any] = field(default_factory=dict)
|
||||
epoch: int = 0
|
||||
iteration: int = 0
|
||||
consumed_samples: int = 0
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
meta: Dict[str, Any] = field(default_factory=dict)
|
||||
config: Dict[str, Any] = field(default_factory=dict)
|
||||
@@ -145,12 +147,9 @@ class Checkpoint:
|
||||
save_path = Path(save_dir)
|
||||
save_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
if get_rank() != 0:
|
||||
return
|
||||
|
||||
meta = {
|
||||
"epoch": self.epoch,
|
||||
"iteration": self.iteration,
|
||||
"consumed_samples": self.consumed_samples,
|
||||
"timestamp": time.strftime("%Y-%m-%dT%H:%M:%S"),
|
||||
**self.meta,
|
||||
}
|
||||
@@ -176,8 +175,9 @@ class Checkpoint:
|
||||
return cls(
|
||||
state_dict=state_dict,
|
||||
epoch=meta.get("epoch", 0),
|
||||
iteration=meta.get("iteration", 0),
|
||||
consumed_samples=meta.get("consumed_samples", 0),
|
||||
extra=extra,
|
||||
meta=meta,
|
||||
config=config,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,123 @@
|
||||
"""Dataset storage serialization helpers (HDF5 / memory-mapped binary)."""
|
||||
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import h5py
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
def save_h5(file_path: str, file_name: str, tensor_group: Dict[str, List[Tensor]]):
|
||||
os.makedirs(file_path, exist_ok=True)
|
||||
full_file_path = os.path.join(file_path, f"{file_name}.h5")
|
||||
with h5py.File(full_file_path, "w") as f:
|
||||
for key, tensors in tensor_group.items():
|
||||
grp = f.create_group(key)
|
||||
for idx, tensor in enumerate(tensors):
|
||||
arr = tensor.cpu().numpy()
|
||||
grp.create_dataset(f"data_{idx}", data=arr)
|
||||
|
||||
|
||||
def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
|
||||
tensor_group: Dict[str, List[Tensor]] = {}
|
||||
|
||||
root_path = Path(file_path)
|
||||
if root_path.is_file() and root_path.suffix in (".h5", ".hdf5"):
|
||||
h5_files = [root_path]
|
||||
else:
|
||||
h5_files = list(root_path.rglob("*.h5")) + list(root_path.rglob("*.hdf5"))
|
||||
|
||||
for h5_file in h5_files:
|
||||
with h5py.File(h5_file, "r") as f:
|
||||
for key in f.keys():
|
||||
grp = f[key]
|
||||
dsets = []
|
||||
for dset_name in grp.keys():
|
||||
dset = grp[dset_name]
|
||||
tensor = torch.from_numpy(dset[:])
|
||||
if share_memory:
|
||||
tensor = tensor.share_memory_()
|
||||
dsets.append(tensor)
|
||||
|
||||
if tensor_group.get(key) is None:
|
||||
tensor_group[key] = []
|
||||
tensor_group[key].extend(dsets)
|
||||
|
||||
return tensor_group
|
||||
|
||||
|
||||
def save_bin(
|
||||
file_path: str,
|
||||
tensor_group: Dict[str, List[Tensor]],
|
||||
record_keys: Optional[List[str]] = None,
|
||||
):
|
||||
"""Save tensors as memory-mapped binary files.
|
||||
|
||||
When *record_keys* is provided, those keys are written with per-record
|
||||
cumulative offsets in ``meta.json`` so that ``MmapStore.fetch_record``
|
||||
can slice individual records from the concatenated binary without
|
||||
cross-record concatenation. Keys not in *record_keys* (e.g. SEQ
|
||||
``sequence``) are written as a single contiguous stream without
|
||||
offsets, preserving backward compatibility.
|
||||
|
||||
Nested keys (``List[List[Tensor]]`` such as GRPO ``responses``) are
|
||||
not supported in bin format — use H5 for those.
|
||||
"""
|
||||
os.makedirs(file_path, exist_ok=True)
|
||||
record_keys = set(record_keys or [])
|
||||
meta = {}
|
||||
for key, tensors in tensor_group.items():
|
||||
if tensors and isinstance(tensors[0], list):
|
||||
raise ValueError(
|
||||
f"Nested key '{key}' (List[List[Tensor]]) is not supported "
|
||||
f"in bin format. Use H5 or JSONL storage instead."
|
||||
)
|
||||
cat = torch.cat(tensors, dim=0)
|
||||
entry: Dict[str, Any] = {
|
||||
"shape": list(cat.shape),
|
||||
"dtype": str(cat.dtype).split(".")[-1],
|
||||
}
|
||||
if key in record_keys:
|
||||
offsets = [0]
|
||||
for t in tensors:
|
||||
offsets.append(offsets[-1] + t.shape[0])
|
||||
entry["offsets"] = offsets
|
||||
meta[key] = entry
|
||||
np.asarray(cat.cpu().numpy()).tofile(os.path.join(file_path, f"{key}.bin"))
|
||||
with open(os.path.join(file_path, "meta.json"), "w") as f:
|
||||
json.dump(meta, f)
|
||||
|
||||
|
||||
def load_bin(file_path: str) -> Dict[str, List[Tensor]]:
|
||||
with open(os.path.join(file_path, "meta.json"), "r") as f:
|
||||
meta = json.load(f)
|
||||
segments: Dict[str, List[Tensor]] = {}
|
||||
for key, info in meta.items():
|
||||
arr = np.memmap(
|
||||
os.path.join(file_path, f"{key}.bin"),
|
||||
dtype=info["dtype"],
|
||||
mode="r",
|
||||
shape=tuple(info["shape"]),
|
||||
)
|
||||
segments[key] = [torch.from_numpy(arr)]
|
||||
return segments
|
||||
|
||||
|
||||
def load_bin_offsets(file_path: str) -> Dict[str, List[int]]:
|
||||
"""Read per-record cumulative offsets from ``meta.json``.
|
||||
|
||||
Returns an empty dict when no key has offsets (legacy bin files),
|
||||
in which case record-mode access falls back to per-record segment
|
||||
indexing (H5/JSONL layout).
|
||||
"""
|
||||
with open(os.path.join(file_path, "meta.json"), "r") as f:
|
||||
meta = json.load(f)
|
||||
offsets: Dict[str, List[int]] = {}
|
||||
for key, info in meta.items():
|
||||
if "offsets" in info:
|
||||
offsets[key] = info["offsets"]
|
||||
return offsets
|
||||
@@ -1,3 +1,4 @@
|
||||
from functools import cached_property
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from jinja2 import Template
|
||||
@@ -29,7 +30,19 @@ class ChatTemplate:
|
||||
self.description = description
|
||||
self.default_variables = default_variables or {}
|
||||
self.special_tokens = special_tokens or {}
|
||||
self._compiled: Template = Template(template_str)
|
||||
|
||||
@cached_property
|
||||
def _compiled(self) -> Template:
|
||||
"""Lazy-compiled Jinja2 template, cached on first access.
|
||||
|
||||
The compiled :class:`~jinja2.Template` holds a dynamically-generated
|
||||
``root`` render function whose ``__module__`` is ``None``; under
|
||||
``pickle`` it falls back to ``__main__`` and breaks ``spawn``-based
|
||||
multiprocessing. By deferring compilation to first access, the
|
||||
default pickle protocol serialises only ``template_str``; each
|
||||
worker rebuilds the cache on first render.
|
||||
"""
|
||||
return Template(self.template_str)
|
||||
|
||||
@classmethod
|
||||
def from_string(
|
||||
|
||||
@@ -164,7 +164,14 @@ class AutoTokenizer:
|
||||
- tokenizer.bos_token → returns string
|
||||
- tokenizer.bos_token_id → returns corresponding integer ID
|
||||
- tokenizer.stop_ids → returns list of corresponding integer IDs for all special tokens
|
||||
|
||||
Internal/private attrs are not intercepted: during unpickle
|
||||
``__dict__`` is empty, so probing ``self._special_token_map``
|
||||
would recurse infinitely.
|
||||
"""
|
||||
if key.startswith("_"):
|
||||
raise AttributeError(key)
|
||||
|
||||
# Handle stop_ids - return IDs for all special tokens
|
||||
if key == "stop_ids":
|
||||
stop_ids = []
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
from astrai.trainer.optim import Muon
|
||||
from astrai.trainer.schedule import BaseScheduler, SchedulerFactory
|
||||
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
|
||||
from astrai.trainer.train_callback import (
|
||||
@@ -10,8 +9,6 @@ from astrai.trainer.trainer import Trainer
|
||||
__all__ = [
|
||||
# Main trainer
|
||||
"Trainer",
|
||||
# Optimizer
|
||||
"Muon",
|
||||
# Strategy factory
|
||||
"StrategyFactory",
|
||||
"BaseStrategy",
|
||||
|
||||
@@ -1,42 +1,25 @@
|
||||
from typing import Any, Callable, Dict
|
||||
from typing import Dict
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
def _grad_stat(
|
||||
model: nn.Module, fn: Callable[[torch.Tensor], Any], default: Any
|
||||
) -> dict:
|
||||
results = {}
|
||||
for name, param in model.named_parameters():
|
||||
results[name] = default
|
||||
if param.grad is not None:
|
||||
results[name] = fn(param.grad.data)
|
||||
return results
|
||||
def grad_norm(model: nn.Module, per_param: bool = False) -> float | Dict[str, float]:
|
||||
grads = [p.grad.detach() for p in model.parameters() if p.grad is not None]
|
||||
if not grads:
|
||||
return 0.0
|
||||
|
||||
|
||||
def grad_norm(model: nn.Module, norm_type: int = 2) -> Dict[str, float]:
|
||||
return _grad_stat(model, lambda g: g.norm(norm_type).item(), 0.0)
|
||||
|
||||
|
||||
def grad_std(model: nn.Module) -> Dict[str, float]:
|
||||
return _grad_stat(model, lambda g: g.std().item(), 0.0)
|
||||
|
||||
|
||||
def grad_max(model: nn.Module) -> Dict[str, float]:
|
||||
return _grad_stat(model, lambda g: g.max().item(), -float("inf"))
|
||||
|
||||
|
||||
def grad_min(model: nn.Module) -> Dict[str, float]:
|
||||
return _grad_stat(model, lambda g: g.min().item(), float("inf"))
|
||||
|
||||
|
||||
def grad_mean(model: nn.Module) -> Dict[str, float]:
|
||||
return _grad_stat(model, lambda g: g.mean().item(), 0.0)
|
||||
|
||||
|
||||
def grad_nan_num(model: nn.Module) -> Dict[str, int]:
|
||||
return _grad_stat(model, lambda g: g.isnan().sum().item(), 0)
|
||||
total_sq = torch.stack([g.pow(2).sum() for g in grads]).sum()
|
||||
if per_param:
|
||||
norms = {}
|
||||
for name, param in model.named_parameters():
|
||||
if param.grad is not None:
|
||||
norms[name] = param.grad.norm(2).item()
|
||||
else:
|
||||
norms[name] = 0.0
|
||||
norms["total"] = total_sq.sqrt().item()
|
||||
return norms
|
||||
return total_sq.sqrt().item()
|
||||
|
||||
|
||||
def ctx_get_loss(ctx):
|
||||
@@ -52,24 +35,4 @@ def ctx_get_val_loss(ctx):
|
||||
|
||||
|
||||
def ctx_get_grad_norm(ctx):
|
||||
return grad_norm(ctx.model)
|
||||
|
||||
|
||||
def ctx_get_grad_std(ctx):
|
||||
return grad_std(ctx.model)
|
||||
|
||||
|
||||
def ctx_get_grad_max(ctx):
|
||||
return grad_max(ctx.model)
|
||||
|
||||
|
||||
def ctx_get_grad_min(ctx):
|
||||
return grad_min(ctx.model)
|
||||
|
||||
|
||||
def ctx_get_grad_mean(ctx):
|
||||
return grad_mean(ctx.model)
|
||||
|
||||
|
||||
def ctx_get_grad_nan_num(ctx):
|
||||
return grad_nan_num(ctx.model)
|
||||
return ctx.grad_norm
|
||||
|
||||
@@ -1,143 +0,0 @@
|
||||
import torch
|
||||
from torch.optim import Optimizer
|
||||
|
||||
|
||||
def _zeropower_via_newtonschulz(G: torch.Tensor, steps: int = 5):
|
||||
assert G.ndim == 2
|
||||
X = G
|
||||
scale = max(1, G.size(0) / G.size(1)) ** 0.5
|
||||
X = X / (X.norm() + 1e-7) * scale
|
||||
if steps == 0:
|
||||
return X
|
||||
a, b, c = (3.4445, -4.7750, 2.0315)
|
||||
for _ in range(steps):
|
||||
A = X @ X.T
|
||||
B = A @ X
|
||||
X = a * X + b * B + c * (A @ B)
|
||||
return X
|
||||
|
||||
|
||||
class Muon(Optimizer):
|
||||
def __init__(
|
||||
self,
|
||||
params,
|
||||
lr: float = 2e-3,
|
||||
momentum: float = 0.95,
|
||||
weight_decay: float = 0.0,
|
||||
nesterov: bool = True,
|
||||
ns_steps: int = 5,
|
||||
adamw_lr: float = None,
|
||||
adamw_betas: tuple = (0.9, 0.95),
|
||||
adamw_eps: float = 1e-8,
|
||||
adamw_wd: float = 0.0,
|
||||
):
|
||||
defaults = dict(
|
||||
lr=lr,
|
||||
momentum=momentum,
|
||||
weight_decay=weight_decay,
|
||||
nesterov=nesterov,
|
||||
ns_steps=ns_steps,
|
||||
adamw_lr=adamw_lr if adamw_lr is not None else lr * 0.1,
|
||||
adamw_betas=adamw_betas,
|
||||
adamw_eps=adamw_eps,
|
||||
adamw_wd=adamw_wd,
|
||||
)
|
||||
super().__init__(params, defaults)
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
loss = None
|
||||
if closure is not None:
|
||||
with torch.enable_grad():
|
||||
loss = closure()
|
||||
|
||||
for group in self.param_groups:
|
||||
params_2d, params_1d = [], []
|
||||
grads_2d, grads_1d = [], []
|
||||
|
||||
for p in group["params"]:
|
||||
if p.grad is None:
|
||||
continue
|
||||
if p.grad.is_sparse:
|
||||
raise RuntimeError("Muon does not support sparse gradients")
|
||||
if p.ndim >= 2:
|
||||
params_2d.append(p)
|
||||
grads_2d.append(p.grad)
|
||||
else:
|
||||
params_1d.append(p)
|
||||
grads_1d.append(p.grad)
|
||||
|
||||
if params_2d:
|
||||
self._muon_update_foreach(params_2d, grads_2d, group)
|
||||
if params_1d:
|
||||
self._adamw_update_foreach(params_1d, grads_1d, group)
|
||||
|
||||
return loss
|
||||
|
||||
def _muon_update_foreach(self, params_2d, grads_2d, group):
|
||||
lr = group["lr"]
|
||||
momentum = group["momentum"]
|
||||
wd = group["weight_decay"]
|
||||
nesterov = group["nesterov"]
|
||||
ns_steps = group["ns_steps"]
|
||||
|
||||
if wd != 0:
|
||||
torch._foreach_mul_(params_2d, 1 - lr * wd)
|
||||
|
||||
if nesterov:
|
||||
grads_2d = torch._foreach_add(grads_2d, params_2d, alpha=wd)
|
||||
|
||||
bufs = []
|
||||
for p, grad in zip(params_2d, grads_2d):
|
||||
state = self.state[p]
|
||||
if "momentum_buffer" not in state:
|
||||
state["momentum_buffer"] = torch.zeros_like(grad)
|
||||
bufs.append(state["momentum_buffer"])
|
||||
|
||||
torch._foreach_lerp_(bufs, grads_2d, 1 - momentum)
|
||||
|
||||
for p, buf in zip(params_2d, bufs):
|
||||
update = _zeropower_via_newtonschulz(buf, steps=ns_steps)
|
||||
scale = max(1, p.size(0) / p.size(1)) ** 0.5
|
||||
p.add_(update, alpha=-lr * scale)
|
||||
|
||||
def _adamw_update_foreach(self, params_1d, grads_1d, group):
|
||||
lr = group["adamw_lr"]
|
||||
betas = group["adamw_betas"]
|
||||
eps = group["adamw_eps"]
|
||||
wd = group["adamw_wd"]
|
||||
|
||||
steps: list[int] = []
|
||||
exp_avgs, exp_avg_sqs = [], []
|
||||
has_state = []
|
||||
for p in params_1d:
|
||||
state = self.state[p]
|
||||
if not state:
|
||||
state["step"] = 0
|
||||
state["exp_avg"] = torch.zeros_like(p)
|
||||
state["exp_avg_sq"] = torch.zeros_like(p)
|
||||
has_state.append(False)
|
||||
else:
|
||||
has_state.append(True)
|
||||
state["step"] += 1
|
||||
steps.append(state["step"])
|
||||
exp_avgs.append(state["exp_avg"])
|
||||
exp_avg_sqs.append(state["exp_avg_sq"])
|
||||
|
||||
beta1, beta2 = betas
|
||||
|
||||
torch._foreach_lerp_(exp_avgs, grads_1d, 1 - beta1)
|
||||
grads_sq = torch._foreach_mul(grads_1d, grads_1d)
|
||||
torch._foreach_lerp_(exp_avg_sqs, grads_sq, 1 - beta2)
|
||||
|
||||
bias_correction1 = [1 - beta1**s for s in steps]
|
||||
bias_correction2 = [1 - beta2**s for s in steps]
|
||||
|
||||
if wd != 0:
|
||||
torch._foreach_mul_(params_1d, 1 - lr * wd)
|
||||
|
||||
exp_avg_corrected = torch._foreach_div(exp_avgs, bias_correction1)
|
||||
denom = torch._foreach_div(exp_avg_sqs, bias_correction2)
|
||||
denom = torch._foreach_sqrt(denom)
|
||||
torch._foreach_add_(denom, eps)
|
||||
torch._foreach_addcdiv_(params_1d, exp_avg_corrected, denom, value=-lr)
|
||||
@@ -53,7 +53,7 @@ class CosineScheduler(BaseScheduler):
|
||||
optimizer,
|
||||
warmup_steps: int,
|
||||
lr_decay_steps: int,
|
||||
min_rate: float = 0.05,
|
||||
min_rate: float = 0.01,
|
||||
last_epoch: int = -1,
|
||||
):
|
||||
self.warmup_steps = warmup_steps
|
||||
@@ -65,11 +65,15 @@ class CosineScheduler(BaseScheduler):
|
||||
def get_lr(self) -> List[float]:
|
||||
# warmup
|
||||
if self.last_epoch < self.warmup_steps:
|
||||
warmup_factor = max(self.min_rate, self.last_epoch / self.warmup_steps)
|
||||
warmup_factor = max(
|
||||
self.min_rate, self.last_epoch / max(self.warmup_steps, 1)
|
||||
)
|
||||
return [base_lr * warmup_factor for base_lr in self.base_lrs]
|
||||
|
||||
# cosine decay
|
||||
decay_progress = (self.last_epoch - self.warmup_steps) / self.lr_decay_steps
|
||||
decay_progress = (self.last_epoch - self.warmup_steps) / max(
|
||||
self.lr_decay_steps, 1
|
||||
)
|
||||
decay_progress = min(decay_progress, 1.0)
|
||||
cosine_decay = 0.5 * (1.0 + math.cos(math.pi * decay_progress))
|
||||
decay_factor = max(self.min_rate, cosine_decay)
|
||||
@@ -104,7 +108,7 @@ class SGDRScheduler(BaseScheduler):
|
||||
optimizer,
|
||||
warmup_steps: int,
|
||||
cycle_length: int,
|
||||
min_rate: float = 0.05,
|
||||
min_rate: float = 0.01,
|
||||
t_mult: int = 2,
|
||||
last_epoch: int = -1,
|
||||
):
|
||||
@@ -118,7 +122,9 @@ class SGDRScheduler(BaseScheduler):
|
||||
def get_lr(self):
|
||||
# warmup
|
||||
if self.last_epoch < self.warmup_steps:
|
||||
warmup_factor = max(self.min_rate, self.last_epoch / self.warmup_steps)
|
||||
warmup_factor = max(
|
||||
self.min_rate, self.last_epoch / max(self.warmup_steps, 1)
|
||||
)
|
||||
return [base_lr * warmup_factor for base_lr in self.base_lrs]
|
||||
|
||||
# SGDR
|
||||
@@ -182,7 +188,7 @@ class WSDScheduler(BaseScheduler):
|
||||
warmup_steps: int,
|
||||
stable_steps: int,
|
||||
decay_steps: int,
|
||||
min_rate: float = 0.0,
|
||||
min_rate: float = 0.01,
|
||||
last_epoch: int = -1,
|
||||
):
|
||||
self.warmup_steps = warmup_steps
|
||||
@@ -194,7 +200,7 @@ class WSDScheduler(BaseScheduler):
|
||||
|
||||
def get_lr(self) -> List[float]:
|
||||
if self.last_epoch < self.warmup_steps:
|
||||
factor = self.last_epoch / max(self.warmup_steps, 1)
|
||||
factor = max(self.min_rate, self.last_epoch / max(self.warmup_steps, 1))
|
||||
return [base_lr * factor for base_lr in self.base_lrs]
|
||||
|
||||
offset = self.last_epoch - self.warmup_steps
|
||||
|
||||
+67
-39
@@ -98,7 +98,6 @@ class BaseStrategy(ABC):
|
||||
self.model = model
|
||||
self.device = device
|
||||
self.executor = kwargs.pop("executor", None)
|
||||
self.model_fn = kwargs.pop("model_fn", None)
|
||||
self.extra_kwargs = kwargs
|
||||
|
||||
@abstractmethod
|
||||
@@ -196,7 +195,7 @@ class SFTStrategy(BaseStrategy):
|
||||
|
||||
ignore_index = -100
|
||||
input_mask = make_doc_boundary_mask(position_ids)
|
||||
target_ids = target_ids.masked_fill(loss_mask == 0, ignore_index)
|
||||
target_ids = target_ids.masked_fill(~loss_mask, ignore_index)
|
||||
logits = self.model(
|
||||
input_ids=input_ids, position_ids=position_ids, input_mask=input_mask
|
||||
)["logits"]
|
||||
@@ -223,14 +222,13 @@ class DPOStrategy(BaseStrategy):
|
||||
self,
|
||||
model: nn.Module,
|
||||
device: str,
|
||||
ref_model: nn.Module,
|
||||
beta: float = 0.1,
|
||||
reduction: str = "mean",
|
||||
reduction: str = "sum",
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(model, device, **kwargs)
|
||||
self.ref_model = create_ref_model(
|
||||
self.model_fn, self.executor.unwrap_model(model)
|
||||
).to(device=self.device)
|
||||
self.ref_model = ref_model
|
||||
self.beta = beta
|
||||
self.reduction = reduction
|
||||
|
||||
@@ -267,42 +265,45 @@ class DPOStrategy(BaseStrategy):
|
||||
class GRPOStrategy(BaseStrategy):
|
||||
"""Group Relative Policy Optimization strategy.
|
||||
|
||||
On-policy GRPO following DeepSeek-R1: the policy model is updated while
|
||||
a frozen ref_model stores the old-policy log-probs. ratio = exp(logπ_θ - logπ_ref),
|
||||
clipped PPO objective. Call ``sync_ref_model()`` after each data-generation round.
|
||||
Implements GRPO following DeepSeek-R1 with token-level PPO clipping.
|
||||
Advantages are group-normalized from scalar per-response rewards and
|
||||
broadcast across all response tokens. The loss is computed **only on
|
||||
response tokens** — prompt tokens are masked out.
|
||||
|
||||
Three model roles are distinguished:
|
||||
|
||||
* **Policy** ``self.model`` — the model being trained.
|
||||
* **Old policy** ``self.old_model`` — the behaviour policy that generated
|
||||
the responses. Used for the importance sampling ratio
|
||||
``ρ = π_θ / π_old``. Synced externally after each data-generation round.
|
||||
* **Reference model** ``self.ref_model`` — a frozen copy of the initial
|
||||
policy (typically the SFT checkpoint) used **only** for the KL
|
||||
regularisation term. It is never updated during training.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: nn.Module,
|
||||
device: str,
|
||||
old_model: nn.Module,
|
||||
ref_model: nn.Module,
|
||||
clip_eps: float = 0.2,
|
||||
kl_coef: float = 0.01,
|
||||
group_size: int = 4,
|
||||
reduction: str = "mean",
|
||||
sync_interval: int = 200,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(model, device, **kwargs)
|
||||
self.ref_model = create_ref_model(
|
||||
self.model_fn, self.executor.unwrap_model(model)
|
||||
).to(device=self.device)
|
||||
self.old_model = old_model
|
||||
self.ref_model = ref_model
|
||||
self.clip_eps = clip_eps
|
||||
self.kl_coef = kl_coef
|
||||
self.group_size = group_size
|
||||
self.reduction = reduction
|
||||
self.sync_interval = sync_interval
|
||||
self._step = 0
|
||||
|
||||
def sync_ref_model(self):
|
||||
"""Copy current model weights to ref model."""
|
||||
self.ref_model.load_state_dict(self.executor.unwrap_model(self.model))
|
||||
def sync_old_model(self):
|
||||
"""Copy current policy weights to old model."""
|
||||
self.old_model.load_state_dict(self.executor.unwrap_model(self.model))
|
||||
|
||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||
self._step += 1
|
||||
if self._step % self.sync_interval == 0:
|
||||
self.sync_ref_model()
|
||||
|
||||
batch = move_to_device(batch, self.device)
|
||||
prompts = batch["prompts"]
|
||||
responses = batch["responses"]
|
||||
@@ -313,33 +314,60 @@ class GRPOStrategy(BaseStrategy):
|
||||
responses_flat = responses.view(-1, response_len)
|
||||
masks_flat = masks.view(-1, response_len)
|
||||
prompt_expanded = prompts.unsqueeze(1).repeat(1, group_size, 1).flatten(0, 1)
|
||||
prompt_len = prompt_expanded.size(1)
|
||||
|
||||
full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-1)
|
||||
full_masks = torch.cat([torch.ones_like(prompt_expanded), masks_flat], dim=-1)
|
||||
|
||||
log_probs_policy = get_logprobs(
|
||||
self.model, full_sequences, full_masks, self.reduction
|
||||
)
|
||||
log_probs_policy = log_probs_policy.view(batch_size, group_size)
|
||||
# Prompt tokens are masked out (0) so logprobs are computed only for
|
||||
# response tokens. get_logprobs shifts the mask by one position, so
|
||||
# the first response token's logprob (predicted from the last prompt
|
||||
# token) is correctly included.
|
||||
full_masks = torch.cat([torch.zeros_like(prompt_expanded), masks_flat], dim=-1)
|
||||
|
||||
# get_logprobs returns [B*G, S-1] (S = prompt_len + response_len).
|
||||
# Response token logprobs occupy the last ``response_len`` positions
|
||||
# (the first response token is predicted from the last prompt token).
|
||||
token_log_probs_policy = get_logprobs(
|
||||
self.model, full_sequences, full_masks, "none"
|
||||
)[:, prompt_len - 1 :]
|
||||
with torch.no_grad():
|
||||
log_probs_ref = get_logprobs(
|
||||
self.ref_model, full_sequences, full_masks, self.reduction
|
||||
)
|
||||
log_probs_ref = log_probs_ref.view(batch_size, group_size)
|
||||
token_log_probs_old = get_logprobs(
|
||||
self.old_model, full_sequences, full_masks, "none"
|
||||
)[:, prompt_len - 1 :]
|
||||
token_log_probs_ref = get_logprobs(
|
||||
self.ref_model, full_sequences, full_masks, "none"
|
||||
)[:, prompt_len - 1 :]
|
||||
|
||||
eps = torch.finfo(log_probs_policy.dtype).eps
|
||||
# Reshape to [B, G, response_len]
|
||||
token_log_probs_policy = token_log_probs_policy.view(batch_size, group_size, -1)
|
||||
token_log_probs_old = token_log_probs_old.view(batch_size, group_size, -1)
|
||||
token_log_probs_ref = token_log_probs_ref.view(batch_size, group_size, -1)
|
||||
token_masks = masks_flat.view(batch_size, group_size, -1).float()
|
||||
|
||||
# Group-normalized advantages from scalar per-response rewards.
|
||||
eps = 1e-8
|
||||
mean = rewards.mean(dim=-1, keepdim=True)
|
||||
std = rewards.std(dim=-1, keepdim=True)
|
||||
std = rewards.std(dim=-1, keepdim=True, unbiased=False)
|
||||
advantages = (rewards - mean) / (std + eps)
|
||||
# Broadcast scalar advantage to every response token: [B, G, 1]
|
||||
advantages = advantages.unsqueeze(-1)
|
||||
|
||||
ratio = torch.exp(log_probs_policy - log_probs_ref)
|
||||
# Token-level ratio (π_θ / π_old) and PPO clipping.
|
||||
log_ratio = token_log_probs_policy - token_log_probs_old
|
||||
ratio = torch.exp(log_ratio)
|
||||
|
||||
surr1 = ratio * advantages
|
||||
surr2 = torch.clamp(ratio, 1 - self.clip_eps, 1 + self.clip_eps) * advantages
|
||||
per_token_policy_loss = -torch.min(surr1, surr2)
|
||||
token_count = token_masks.sum().clamp(min=1.0)
|
||||
policy_loss = (per_token_policy_loss * token_masks).sum() / token_count
|
||||
|
||||
# KL penalty to frozen reference model with k1 estimator (non-negative):
|
||||
# k1 = π_ref / π_θ - log(π_ref / π_θ) - 1, where π_ref / π_θ = exp(log_ref - log_policy).
|
||||
log_ref_ratio = token_log_probs_ref - token_log_probs_policy
|
||||
r = torch.exp(log_ref_ratio)
|
||||
kl_per_token = r - torch.log(r + eps) - 1.0
|
||||
kl_penalty = self.kl_coef * (kl_per_token * token_masks).sum() / token_count
|
||||
|
||||
policy_loss = -torch.min(surr1, surr2).mean()
|
||||
kl_penalty = self.kl_coef * (log_probs_policy - log_probs_ref).square().mean()
|
||||
total_loss = policy_loss + kl_penalty
|
||||
|
||||
return total_loss
|
||||
|
||||
+104
-100
@@ -9,21 +9,15 @@ from typing import IO, Callable, List, Optional, Protocol, runtime_checkable
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
from torch.nn.utils import clip_grad_norm_
|
||||
from torch.utils.checkpoint import checkpoint as torch_checkpoint
|
||||
from tqdm import tqdm
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.parallel import only_on_rank
|
||||
from astrai.parallel.setup import get_current_device, get_rank
|
||||
from astrai.parallel.setup import get_current_device
|
||||
from astrai.serialization import Checkpoint
|
||||
from astrai.trainer.metric_util import (
|
||||
ctx_get_grad_max,
|
||||
ctx_get_grad_mean,
|
||||
ctx_get_grad_min,
|
||||
ctx_get_grad_nan_num,
|
||||
ctx_get_grad_norm,
|
||||
ctx_get_grad_std,
|
||||
ctx_get_loss,
|
||||
ctx_get_lr,
|
||||
ctx_get_val_loss,
|
||||
@@ -86,7 +80,9 @@ class GradientClippingCallback(TrainCallback):
|
||||
self.max_grad_norm = max_grad_norm
|
||||
|
||||
def on_optimizer_step(self, context: TrainContext):
|
||||
clip_grad_norm_(context.model.parameters(), self.max_grad_norm)
|
||||
context.grad_norm = context.executor.clip_grad_norm(
|
||||
context.model, self.max_grad_norm
|
||||
)
|
||||
|
||||
|
||||
@CallbackFactory.register("gradient_checkpointing")
|
||||
@@ -143,34 +139,38 @@ class CheckpointCallback(TrainCallback):
|
||||
self.interval = interval
|
||||
self.weight_only = weight_only
|
||||
self.save_extra_fn = save_extra_fn or CheckpointCallback.save_extra
|
||||
self.last_ckpt_iter = 0
|
||||
self.last_ckpt_step = None
|
||||
|
||||
def on_train_begin(self, context: TrainContext):
|
||||
self.last_ckpt_step = context.optimizer_step
|
||||
|
||||
def _save_checkpoint(self, context: TrainContext):
|
||||
state_dict = context.executor.unwrap_model(context.model)
|
||||
self.last_ckpt_iter = context.iteration
|
||||
self.last_ckpt_step = context.optimizer_step
|
||||
|
||||
if get_rank() == 0:
|
||||
save_path = os.path.join(
|
||||
self.save_dir, f"epoch_{context.epoch}_iter_{context.iteration}"
|
||||
)
|
||||
extra = self.save_extra_fn(context)
|
||||
meta = context.config.to_dict()
|
||||
context.checkpoint = Checkpoint(
|
||||
state_dict=state_dict,
|
||||
epoch=context.epoch,
|
||||
iteration=context.iteration,
|
||||
extra=extra,
|
||||
meta=meta,
|
||||
config=context.model_config,
|
||||
)
|
||||
context.checkpoint.save(save_path)
|
||||
with context.executor.checkpoint_context(context.model) as state_dict:
|
||||
if state_dict is not None:
|
||||
save_path = os.path.join(
|
||||
self.save_dir,
|
||||
f"epoch_{context.epoch}_step_{context.optimizer_step}",
|
||||
)
|
||||
extra = self.save_extra_fn(context)
|
||||
meta = context.config.to_dict()
|
||||
context.checkpoint = Checkpoint(
|
||||
state_dict=state_dict,
|
||||
epoch=context.epoch,
|
||||
consumed_samples=context.consumed_samples,
|
||||
config=context.model_config,
|
||||
extra=extra,
|
||||
meta=meta,
|
||||
)
|
||||
context.checkpoint.save(save_path)
|
||||
|
||||
def on_batch_end(self, context: TrainContext):
|
||||
if context.iteration - self.last_ckpt_iter >= self.interval:
|
||||
if context.optimizer_step - self.last_ckpt_step >= self.interval:
|
||||
self._save_checkpoint(context)
|
||||
|
||||
def on_train_end(self, context: TrainContext):
|
||||
if context.iteration != self.last_ckpt_iter:
|
||||
if context.optimizer_step != self.last_ckpt_step:
|
||||
self._save_checkpoint(context)
|
||||
|
||||
def on_error(self, context: TrainContext):
|
||||
@@ -202,19 +202,23 @@ class ProgressBarCallback(TrainCallback):
|
||||
|
||||
@only_on_rank(0)
|
||||
def on_epoch_begin(self, context: TrainContext):
|
||||
total_steps = len(context.dataloader) // context.executor.grad_accum_steps
|
||||
self.progress_bar = tqdm(
|
||||
context.dataloader,
|
||||
total=total_steps,
|
||||
desc=f"Epoch {context.epoch + 1}/{self.num_epoch}",
|
||||
dynamic_ncols=True,
|
||||
file=self.file or sys.stdout,
|
||||
)
|
||||
|
||||
@only_on_rank(0)
|
||||
def on_batch_end(self, context: TrainContext):
|
||||
def on_optimizer_step(self, context: TrainContext):
|
||||
postfix = {
|
||||
"step": f"{context.optimizer_step:d}",
|
||||
"loss": f"{context.loss:.4f}",
|
||||
"lr": f"{context.optimizer.param_groups[-1]['lr']:.2e}",
|
||||
}
|
||||
if context.grad_norm is not None:
|
||||
postfix["grad_norm"] = f"{context.grad_norm:.2f}"
|
||||
if context.val_loss is not None:
|
||||
postfix["val_loss"] = f"{context.val_loss:.4f}"
|
||||
self.progress_bar.set_postfix(postfix)
|
||||
@@ -227,19 +231,20 @@ class ProgressBarCallback(TrainCallback):
|
||||
self.progress_bar.close()
|
||||
|
||||
|
||||
@CallbackFactory.register("metric_logger")
|
||||
class MetricLoggerCallback(TrainCallback):
|
||||
@CallbackFactory.register("metric")
|
||||
class MetricCallback(TrainCallback):
|
||||
def __init__(
|
||||
self,
|
||||
log_dir: str,
|
||||
save_interval: int,
|
||||
log_interval: int = 10,
|
||||
metrics: List[str] = None,
|
||||
val_step: int = 0,
|
||||
):
|
||||
self.last_log_iter = 0
|
||||
self.last_log_flush_step = None
|
||||
self.save_interval = save_interval
|
||||
self.log_interval = log_interval
|
||||
self.metrics = metrics or ["loss", "lr"]
|
||||
self.val_step = val_step
|
||||
self._next_val_step = 0
|
||||
|
||||
self.log_dir = Path(log_dir) if log_dir else Path.cwd() / "logs"
|
||||
self.log_dir.mkdir(parents=True, exist_ok=True)
|
||||
@@ -251,58 +256,28 @@ class MetricLoggerCallback(TrainCallback):
|
||||
"lr": ctx_get_lr,
|
||||
"val_loss": ctx_get_val_loss,
|
||||
"grad_norm": ctx_get_grad_norm,
|
||||
"grad_std": ctx_get_grad_std,
|
||||
"grad_max": ctx_get_grad_max,
|
||||
"grad_min": ctx_get_grad_min,
|
||||
"grad_mean": ctx_get_grad_mean,
|
||||
"grad_nan_num": ctx_get_grad_nan_num,
|
||||
}
|
||||
|
||||
def _get_log_data(self, context: TrainContext):
|
||||
data = {
|
||||
def _metrics(self, context: TrainContext, names):
|
||||
return {
|
||||
m: self._metric_funcs[m](context)
|
||||
for m in names
|
||||
if self._metric_funcs[m](context) is not None
|
||||
}
|
||||
|
||||
@only_on_rank(0)
|
||||
def _append(self, event_type: str, context: TrainContext, **extra):
|
||||
entry = {
|
||||
"type": event_type,
|
||||
"timestamp": time.strftime("%Y-%m-%dT%H:%M:%S"),
|
||||
"epoch": context.epoch,
|
||||
"iter": context.iteration,
|
||||
"step": context.optimizer_step,
|
||||
"consumed_samples": context.consumed_samples,
|
||||
**extra,
|
||||
}
|
||||
for m in self.metrics:
|
||||
val = self._metric_funcs[m](context)
|
||||
if val is not None:
|
||||
data[m] = val
|
||||
return data
|
||||
self.log_cache.append(entry)
|
||||
|
||||
@only_on_rank(0)
|
||||
def _add_log(self, log_data):
|
||||
self.log_cache.append(log_data)
|
||||
|
||||
@only_on_rank(0)
|
||||
def _save_log(self, epoch, iter):
|
||||
log_file = self.log_dir / f"epoch_{epoch}_iter_{iter}_metric.jsonl"
|
||||
log_file.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
with open(log_file, "w") as f:
|
||||
for log in self.log_cache:
|
||||
f.write(json.dumps(log) + "\n")
|
||||
|
||||
def on_batch_end(self, context):
|
||||
if context.iteration % self.log_interval == 0:
|
||||
log_data = self._get_log_data(context)
|
||||
self._add_log(log_data)
|
||||
|
||||
if context.iteration - self.last_log_iter >= self.save_interval:
|
||||
self._save_log(context.epoch, context.iteration)
|
||||
self.last_log_iter = context.iteration
|
||||
|
||||
def on_train_end(self, context):
|
||||
if context.iteration != self.last_log_iter:
|
||||
self._save_log(context.epoch, context.iteration)
|
||||
|
||||
def on_error(self, context):
|
||||
self._save_log(context.epoch, context.iteration)
|
||||
|
||||
|
||||
@CallbackFactory.register("validation")
|
||||
class ValidationCallback(TrainCallback):
|
||||
def _run_validation(self, context: TrainContext):
|
||||
def _run_validation(self, context: TrainContext) -> float:
|
||||
context.model.eval()
|
||||
|
||||
total_loss = 0.0
|
||||
@@ -314,27 +289,56 @@ class ValidationCallback(TrainCallback):
|
||||
total_loss += loss.item()
|
||||
num_batches += 1
|
||||
|
||||
avg_loss = total_loss / max(num_batches, 1)
|
||||
|
||||
if context.world_size > 1 and dist.is_initialized():
|
||||
loss_tensor = torch.tensor([avg_loss], device=get_current_device())
|
||||
dist.all_reduce(loss_tensor, op=dist.ReduceOp.AVG)
|
||||
avg_loss = loss_tensor.item()
|
||||
stats = torch.tensor(
|
||||
[total_loss, float(num_batches)], device=get_current_device()
|
||||
)
|
||||
dist.all_reduce(stats, op=dist.ReduceOp.SUM)
|
||||
avg_loss = (stats[0] / stats[1]).item()
|
||||
else:
|
||||
avg_loss = total_loss / max(num_batches, 1)
|
||||
|
||||
context.val_loss = avg_loss
|
||||
context.model.train()
|
||||
return avg_loss
|
||||
|
||||
step_count = context.iteration // context.config.grad_accum_steps
|
||||
logger.info(
|
||||
f"Epoch {context.epoch + 1}, Step {step_count}, Val Loss: {avg_loss:.4f}"
|
||||
)
|
||||
def on_train_begin(self, context: TrainContext):
|
||||
self.last_log_flush_step = context.optimizer_step
|
||||
|
||||
def on_optimizer_step(self, context: TrainContext):
|
||||
if context.val_dataloader is None:
|
||||
return
|
||||
cfg = context.config
|
||||
if cfg.val_step <= 0:
|
||||
return
|
||||
step_count = context.iteration // cfg.grad_accum_steps
|
||||
if step_count % cfg.val_step == 0:
|
||||
self._run_validation(context)
|
||||
@only_on_rank(0)
|
||||
def _flush(self, epoch, step):
|
||||
log_file = self.log_dir / f"epoch_{epoch}_step_{step}_metric.jsonl"
|
||||
log_file.parent.mkdir(parents=True, exist_ok=True)
|
||||
with open(log_file, "w") as f:
|
||||
for log in self.log_cache:
|
||||
f.write(json.dumps(log) + "\n")
|
||||
|
||||
def on_optimizer_step(self, context):
|
||||
if (
|
||||
context.val_dataloader is not None
|
||||
and self.val_step > 0
|
||||
and context.optimizer_step >= self._next_val_step
|
||||
):
|
||||
context.val_loss = self._run_validation(context)
|
||||
self._next_val_step = context.optimizer_step + self.val_step
|
||||
self._append("validation", context, val_loss=context.val_loss)
|
||||
|
||||
step_metrics = [m for m in self.metrics if m != "val_loss"]
|
||||
self._append("step", context, **self._metrics(context, step_metrics))
|
||||
|
||||
if context.optimizer_step - self.last_log_flush_step >= self.save_interval:
|
||||
self._flush(context.epoch, context.optimizer_step)
|
||||
self.last_log_flush_step = context.optimizer_step
|
||||
|
||||
def on_epoch_end(self, context):
|
||||
self._append("epoch", context)
|
||||
|
||||
def on_train_end(self, context):
|
||||
if (
|
||||
self.last_log_flush_step is None
|
||||
or context.optimizer_step != self.last_log_flush_step
|
||||
):
|
||||
self._flush(context.epoch, context.optimizer_step)
|
||||
self.last_log_flush_step = context.optimizer_step
|
||||
|
||||
def on_error(self, context):
|
||||
self._flush(context.epoch, context.optimizer_step)
|
||||
|
||||
@@ -7,13 +7,13 @@ import torch.nn as nn
|
||||
from torch.utils.data import DataLoader, random_split
|
||||
|
||||
from astrai.config.train_config import TrainConfig
|
||||
from astrai.dataset import ResumableDistributedSampler
|
||||
from astrai.dataset import RDSampler
|
||||
from astrai.model.components.lora import inject_lora
|
||||
from astrai.parallel.executor import BaseExecutor, ExecutorFactory
|
||||
from astrai.parallel.setup import get_current_device, get_rank, get_world_size
|
||||
from astrai.protocols import OptimizerProtocol, SchedulerProtocol
|
||||
from astrai.serialization import Checkpoint, load_json
|
||||
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
|
||||
from astrai.trainer.strategy import BaseStrategy, StrategyFactory, create_ref_model
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -29,8 +29,9 @@ class TrainContext:
|
||||
executor: BaseExecutor = field(default=None)
|
||||
|
||||
epoch: int = field(default=0)
|
||||
iteration: int = field(default=0)
|
||||
consumed_samples: int = field(default=0)
|
||||
loss: float = field(default=0.0)
|
||||
grad_norm: Optional[float] = field(default=None)
|
||||
val_dataloader: Optional[DataLoader] = field(default=None)
|
||||
val_loss: Optional[float] = field(default=None)
|
||||
|
||||
@@ -38,6 +39,14 @@ class TrainContext:
|
||||
rank: int = field(default=0)
|
||||
kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
@property
|
||||
def optimizer_step(self) -> int:
|
||||
return self.consumed_samples // (
|
||||
self.config.batch_per_device
|
||||
* self.world_size
|
||||
* self.config.grad_accum_steps
|
||||
)
|
||||
|
||||
|
||||
class TrainContextBuilder:
|
||||
def __init__(
|
||||
@@ -45,10 +54,12 @@ class TrainContextBuilder:
|
||||
config: TrainConfig,
|
||||
):
|
||||
self.config = config
|
||||
self._resume_dir: Optional[str] = None
|
||||
self._param_path: Optional[str] = None
|
||||
self._resume: bool = False
|
||||
|
||||
def with_resume_dir(self, resume_dir: Optional[str]) -> Self:
|
||||
self._resume_dir = resume_dir
|
||||
def with_param_path(self, param_path: Optional[str], resume: bool = False) -> Self:
|
||||
self._param_path = param_path
|
||||
self._resume = resume
|
||||
return self
|
||||
|
||||
def build(self) -> TrainContext:
|
||||
@@ -63,11 +74,10 @@ class TrainContextBuilder:
|
||||
|
||||
model = cfg.model_fn()
|
||||
model = model.to(device=device)
|
||||
model.embed_tokens.neftune_noise_alpha = cfg.neftune_alpha
|
||||
|
||||
model_config = {}
|
||||
if self._resume_dir:
|
||||
config_path = Path(self._resume_dir) / "config.json"
|
||||
if self._param_path:
|
||||
config_path = Path(self._param_path) / "config.json"
|
||||
if config_path.exists():
|
||||
model_config = load_json(config_path)
|
||||
|
||||
@@ -83,15 +93,29 @@ class TrainContextBuilder:
|
||||
executor=executor,
|
||||
)
|
||||
|
||||
if self._resume_dir:
|
||||
checkpoint = Checkpoint.load_any(self._resume_dir)
|
||||
if self._param_path:
|
||||
checkpoint = Checkpoint.load_any(self._param_path)
|
||||
if checkpoint is not None:
|
||||
model.load_state_dict(checkpoint.state_dict, strict=False)
|
||||
if checkpoint.config:
|
||||
context.model_config = checkpoint.config
|
||||
context.epoch = checkpoint.epoch or cfg.start_epoch
|
||||
context.iteration = checkpoint.iteration or cfg.start_batch
|
||||
context.checkpoint = checkpoint
|
||||
|
||||
if self._resume:
|
||||
context.epoch = checkpoint.epoch or cfg.start_epoch
|
||||
if checkpoint.consumed_samples > 0:
|
||||
per_step = (
|
||||
cfg.batch_per_device
|
||||
* context.world_size
|
||||
* cfg.grad_accum_steps
|
||||
)
|
||||
context.consumed_samples = (
|
||||
checkpoint.consumed_samples // per_step
|
||||
) * per_step
|
||||
else:
|
||||
context.consumed_samples = (
|
||||
cfg.start_samples * context.world_size
|
||||
)
|
||||
context.checkpoint = checkpoint
|
||||
|
||||
if cfg.lora is not None:
|
||||
inject_lora(
|
||||
@@ -116,8 +140,8 @@ class TrainContextBuilder:
|
||||
cfg.dataset, [n_train, n_val], generator=generator
|
||||
)
|
||||
|
||||
sampler_offset = context.iteration * cfg.batch_per_device
|
||||
sampler = ResumableDistributedSampler(
|
||||
sampler_offset = context.consumed_samples // context.world_size
|
||||
sampler = RDSampler(
|
||||
data_source=train_dataset,
|
||||
start_epoch=context.epoch,
|
||||
start_iter=sampler_offset,
|
||||
@@ -130,10 +154,11 @@ class TrainContextBuilder:
|
||||
num_workers=cfg.num_workers,
|
||||
pin_memory=cfg.pin_memory,
|
||||
prefetch_factor=cfg.prefetch_factor,
|
||||
collate_fn=cfg.collate_fn,
|
||||
)
|
||||
|
||||
if val_dataset is not None:
|
||||
val_sampler = ResumableDistributedSampler(
|
||||
val_sampler = RDSampler(
|
||||
data_source=val_dataset,
|
||||
start_epoch=0,
|
||||
start_iter=0,
|
||||
@@ -147,6 +172,7 @@ class TrainContextBuilder:
|
||||
num_workers=cfg.num_workers,
|
||||
pin_memory=cfg.pin_memory,
|
||||
prefetch_factor=cfg.prefetch_factor,
|
||||
collate_fn=cfg.collate_fn,
|
||||
)
|
||||
|
||||
context.model, context.optimizer, context.dataloader, context.scheduler = (
|
||||
@@ -166,13 +192,26 @@ class TrainContextBuilder:
|
||||
if obj is not None:
|
||||
obj.load_state_dict(extra[name])
|
||||
|
||||
strategy_kwargs = dict(cfg.extra_kwargs)
|
||||
|
||||
if cfg.strategy in ("dpo", "grpo"):
|
||||
ref_model = create_ref_model(
|
||||
cfg.model_fn, executor.unwrap_model(context.model)
|
||||
).to(device=device)
|
||||
strategy_kwargs["ref_model"] = ref_model
|
||||
|
||||
if cfg.strategy == "grpo":
|
||||
old_model = create_ref_model(
|
||||
cfg.model_fn, executor.unwrap_model(context.model)
|
||||
).to(device=device)
|
||||
strategy_kwargs["old_model"] = old_model
|
||||
|
||||
context.strategy = StrategyFactory.create(
|
||||
cfg.strategy,
|
||||
model=context.model,
|
||||
device=device,
|
||||
executor=executor,
|
||||
model_fn=cfg.model_fn,
|
||||
**cfg.extra_kwargs,
|
||||
**strategy_kwargs,
|
||||
)
|
||||
|
||||
return context
|
||||
|
||||
@@ -35,15 +35,14 @@ class Trainer:
|
||||
cfg.ckpt_interval,
|
||||
),
|
||||
CallbackFactory.create(
|
||||
"metric_logger",
|
||||
"metric",
|
||||
log_dir=cfg.log_dir,
|
||||
save_interval=cfg.ckpt_interval,
|
||||
log_interval=cfg.log_interval,
|
||||
metrics=cfg.metrics,
|
||||
val_step=cfg.val_step,
|
||||
),
|
||||
CallbackFactory.create("progress_bar", cfg.n_epoch),
|
||||
CallbackFactory.create("gradient_clipping", cfg.max_grad_norm),
|
||||
CallbackFactory.create("validation"),
|
||||
]
|
||||
return callbacks
|
||||
|
||||
@@ -53,9 +52,11 @@ class Trainer:
|
||||
if method:
|
||||
method(context)
|
||||
|
||||
def _trainer_loop(self, resume_dir: Optional[str] = None):
|
||||
def _trainer_loop(self, param_path: Optional[str] = None, resume: bool = False):
|
||||
context = (
|
||||
TrainContextBuilder(self.train_config).with_resume_dir(resume_dir).build()
|
||||
TrainContextBuilder(self.train_config)
|
||||
.with_param_path(param_path, resume=resume)
|
||||
.build()
|
||||
)
|
||||
executor = context.executor
|
||||
self._call_callbacks("on_train_begin", context)
|
||||
@@ -74,7 +75,9 @@ class Trainer:
|
||||
context.loss = loss.item()
|
||||
stand_loss = loss / executor.grad_accum_steps
|
||||
executor.backward(stand_loss)
|
||||
context.iteration += 1
|
||||
context.consumed_samples += (
|
||||
context.config.batch_per_device * context.world_size
|
||||
)
|
||||
self._call_callbacks("on_batch_end", context)
|
||||
|
||||
if executor.sync_gradients:
|
||||
@@ -94,7 +97,7 @@ class Trainer:
|
||||
finally:
|
||||
self._call_callbacks("on_train_end", context)
|
||||
|
||||
def train(self, resume_dir: Optional[str] = None):
|
||||
def train(self, param_path: Optional[str] = None, resume: bool = False):
|
||||
cfg = self.train_config
|
||||
spawn_parallel_fn(
|
||||
self._trainer_loop,
|
||||
@@ -104,5 +107,6 @@ class Trainer:
|
||||
master_port=cfg.master_port,
|
||||
device_type=cfg.device_type,
|
||||
start_method=cfg.start_method,
|
||||
resume_dir=resume_dir,
|
||||
param_path=param_path,
|
||||
resume=resume,
|
||||
)
|
||||
|
||||
@@ -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,68 @@
|
||||
#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 causal_offset; // -1 = non-causal; >=0 = absolute position of first Q token
|
||||
int num_splits;
|
||||
float scale;
|
||||
|
||||
// Q strides (element offsets for each dim — layout-agnostic)
|
||||
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
|
||||
// KV strides (K and V share the same layout — only base pointers differ)
|
||||
int kv_stride_b, kv_stride_h, kv_stride_l, kv_stride_d;
|
||||
|
||||
// Mask: 2D [batch, kv_len] (mask_q_stride=0) or 3D [batch, q_len, kv_len]
|
||||
int mask_b_stride; // = kv_len (both 2D and 3D)
|
||||
int mask_q_stride; // 2D: 0 (all q rows share); 3D: kv_len
|
||||
|
||||
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 causal_offset;
|
||||
float scale;
|
||||
|
||||
int num_splits;
|
||||
int page_size;
|
||||
int max_pages;
|
||||
|
||||
// Q strides (layout-agnostic)
|
||||
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
|
||||
|
||||
// Mask strides (2D or 3D)
|
||||
int mask_b_stride;
|
||||
int mask_q_stride;
|
||||
|
||||
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,82 @@
|
||||
#include "attn_decode_split_kv.cuh"
|
||||
#include "attn_entry_utils.cuh"
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
#include "attn_decode_split_kv_mma.cuh"
|
||||
#endif
|
||||
|
||||
// Scalar fallback: one warp per query head, split-KV across grid.z.
|
||||
static void launch_scalar_decode(AttentionParams<bf16>& p) {
|
||||
int group_size = p.q_head / p.kv_head;
|
||||
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||
alloc_split_partials(p);
|
||||
|
||||
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
||||
attn_decode_split_kv_kernel<<<dim3(p.batch * p.kv_head, 1, p.num_splits), dim3(32, group_size), smem>>>(p);
|
||||
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
// MMA head-packing requires G <= 16 (BR=16 rows). sm_80+ tensor-core
|
||||
// + cp.async wins even at G=1 (decode is memory-bound, not compute-bound).
|
||||
// STAGES=2 (double-buffer) for D<=128 (smem 16 KB); STAGES=1 for D=256
|
||||
// (double-buffer would be 32 KB, near the 48 KB static cap — keep single
|
||||
// to preserve occupancy).
|
||||
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
|
||||
static void launch_mma_decode(AttentionParams<bf16>& p) {
|
||||
int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
||||
alloc_split_partials(p);
|
||||
|
||||
attn_decode_split_kv_mma_kernel<HEAD_DIM, BC, STAGES><<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
|
||||
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
#endif
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static void dispatch_decode(AttentionParams<bf16>& p) {
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
int G = p.q_head / p.kv_head;
|
||||
if (G >= 1 && G <= 16) {
|
||||
launch_mma_decode<HEAD_DIM, 32>(p);
|
||||
return;
|
||||
}
|
||||
#endif
|
||||
launch_scalar_decode(p);
|
||||
}
|
||||
|
||||
torch::Tensor attn_decode(
|
||||
torch::Tensor q,
|
||||
torch::Tensor k,
|
||||
torch::Tensor v,
|
||||
c10::optional<torch::Tensor> mask,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
int64_t layout
|
||||
) {
|
||||
AttentionParams<bf16> p;
|
||||
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
|
||||
TORCH_CHECK(p.q_len == 1, "Q seq_len must be 1");
|
||||
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
|
||||
|
||||
// O matches Q's original layout
|
||||
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
||||
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
|
||||
p.o = (bf16*)O_view.data_ptr();
|
||||
|
||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p);
|
||||
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("causal_offset") = -1,
|
||||
py::arg("scale") = 0.0,
|
||||
py::arg("layout") = 0,
|
||||
"GQA decode (tensor-core head-packing on sm_80+, scalar fallback)");
|
||||
}
|
||||
@@ -0,0 +1,132 @@
|
||||
#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;
|
||||
|
||||
// Q: [batch, q_head, q_len=1, head_dim] — stride-based
|
||||
float q_reg[8];
|
||||
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
||||
+ lane * hd_per_thread * p.q_stride_d;
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
q_reg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
|
||||
|
||||
// KV: [batch, kv_head, kv_len, head_dim] — stride-based base
|
||||
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||
int mask_base = batch * p.mask_b_stride;
|
||||
|
||||
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
|
||||
|
||||
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);
|
||||
|
||||
// Load K into shared memory (gather from strided global)
|
||||
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 kv_idx = chunk_start + s;
|
||||
int g_off = kv_base + kv_idx * p.kv_stride_l + d_dim * p.kv_stride_d;
|
||||
k_smem[i] = p.k[g_off];
|
||||
}
|
||||
__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;
|
||||
|
||||
int kv_idx = chunk_start + s;
|
||||
if (p.use_mask && p.mask && !p.mask[mask_base + kv_idx])
|
||||
partial = -FLT_MAX;
|
||||
if (p.causal_offset >= 0 && kv_idx > p.causal_offset)
|
||||
partial = -FLT_MAX;
|
||||
|
||||
float new_m = fmaxf(m, partial);
|
||||
float alpha = expf(m - new_m);
|
||||
float beta = expf(partial - new_m);
|
||||
d = d * alpha + beta;
|
||||
|
||||
// V: stride-based read
|
||||
int v_off = kv_base + kv_idx * p.kv_stride_l + lane * hd_per_thread * p.kv_stride_d;
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
acc_reg[i] = acc_reg[i] * alpha + __bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta;
|
||||
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;
|
||||
|
||||
int batch = bh / p.q_head;
|
||||
int q_head = bh % p.q_head;
|
||||
|
||||
size_t split_base = (size_t)bh * p.num_splits;
|
||||
const float* mlp = p.ml_part + split_base * 2;
|
||||
const float* op = p.o_part + split_base * p.head_dim;
|
||||
|
||||
float m = -FLT_MAX, l = 0.0f, acc = 0.0f;
|
||||
for (int s = 0; s < p.num_splits; s++) {
|
||||
float mi = mlp[s * 2];
|
||||
if (mi <= -FLT_MAX) continue;
|
||||
float li = mlp[s * 2 + 1];
|
||||
float nm = fmaxf(m, mi);
|
||||
float corr = __expf(m - nm);
|
||||
float e = __expf(mi - nm);
|
||||
acc = acc * corr + op[s * p.head_dim + d] * e;
|
||||
l = l * corr + li * e;
|
||||
m = nm;
|
||||
}
|
||||
|
||||
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
||||
// Stride-based output write (q_len=1 for decode, so stride_l not needed)
|
||||
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h + d * p.q_stride_d;
|
||||
p.o[o_off] = __float2bfloat16(acc * inv);
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
#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.
|
||||
|
||||
template <int HEAD_DIM, int BC, int STAGES = 2>
|
||||
__global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
||||
constexpr int KD = HEAD_DIM / 16;
|
||||
constexpr int NC8 = BC / 8;
|
||||
constexpr int KT2 = BC / 16;
|
||||
constexpr int DN8 = HEAD_DIM / 8;
|
||||
constexpr int LD = HEAD_DIM;
|
||||
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
|
||||
constexpr int VEC = 8;
|
||||
constexpr int TOTAL = BC * HEAD_DIM;
|
||||
|
||||
const int lane = threadIdx.x;
|
||||
const int 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 ----
|
||||
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
|
||||
const int qra = gid;
|
||||
const int qrb = gid + 8;
|
||||
const bool va = qra < G, vb = qrb < G;
|
||||
unsigned Qa[KD][4];
|
||||
load_q_mma_frags<KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
|
||||
qra, qrb, va, vb, tid4, Qa);
|
||||
|
||||
float Oacc[DN8][4];
|
||||
#pragma unroll
|
||||
for (int j = 0; j < DN8; j++)
|
||||
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
||||
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
||||
|
||||
// KV: stride-based base — [batch, kv_head, kv_len, head_dim]
|
||||
const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||
const int 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);
|
||||
// KV stride-based: contiguous within head_dim (stride_d == 1 typically)
|
||||
int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d;
|
||||
cp_async_16_pred(&dK[off], &p.k[g_off], valid);
|
||||
cp_async_16_pred(&dV[off], &p.v[g_off], 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;
|
||||
|
||||
// Decode: q_len=1, so qrow0=qrow1=0, mask_q_stride irrelevant
|
||||
int maxc = (p.causal_offset >= 0) ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
||||
mma_softmax_tile<NC8, DN8>(kv0, maxc, maxc,
|
||||
0, 0,
|
||||
p.mask_b_stride, 0,
|
||||
batch,
|
||||
p.mask, has_mask,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
|
||||
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc);
|
||||
__syncwarp();
|
||||
|
||||
if constexpr (STAGES == 1) {
|
||||
if (ti + 1 < ti_end)
|
||||
load_tile(ti + 1, 0);
|
||||
}
|
||||
}
|
||||
|
||||
// ---- write UN-normalised partials for this split ----
|
||||
auto split_slot = [&](int h) -> size_t {
|
||||
size_t bh = (size_t)batch * p.q_head + h;
|
||||
return bh * p.num_splits + split;
|
||||
};
|
||||
#pragma unroll
|
||||
for (int dn8 = 0; dn8 < DN8; dn8++) {
|
||||
int d = dn8 * 8 + 2 * tid4;
|
||||
int r0 = gid, r1 = gid + 8;
|
||||
if (r0 < G) {
|
||||
int h = q_head0 + r0;
|
||||
float* op = p.o_part + split_slot(h) * HEAD_DIM;
|
||||
op[d] = Oacc[dn8][0];
|
||||
op[d + 1] = Oacc[dn8][1];
|
||||
}
|
||||
if (r1 < G) {
|
||||
int h = q_head0 + r1;
|
||||
float* op = p.o_part + split_slot(h) * HEAD_DIM;
|
||||
op[d] = Oacc[dn8][2];
|
||||
op[d + 1] = Oacc[dn8][3];
|
||||
}
|
||||
}
|
||||
if (tid4 == 0) {
|
||||
int r0 = gid, r1 = gid + 8;
|
||||
if (r0 < G) {
|
||||
int h = q_head0 + r0;
|
||||
float* mp = p.ml_part + split_slot(h) * 2;
|
||||
mp[0] = m0; mp[1] = l0;
|
||||
}
|
||||
if (r1 < G) {
|
||||
int h = q_head0 + r1;
|
||||
float* mp = p.ml_part + split_slot(h) * 2;
|
||||
mp[0] = m1; mp[1] = l1;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,177 @@
|
||||
#pragma once
|
||||
#include <torch/extension.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include "attn_common.h"
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
|
||||
inline int compute_num_splits(int base_blocks, int tiles_total) {
|
||||
int sm_count = 0;
|
||||
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
|
||||
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
|
||||
return std::max(1, std::min(n, std::min(tiles_total, 32)));
|
||||
}
|
||||
|
||||
// Dispatch head_dim: shared macro — avoids C++20 lambda template syntax.
|
||||
// Usage: DISPATCH_HEAD_DIM(hd, fn, arg)
|
||||
// Expands to: fn<32>(arg); fn<64>(arg); etc.
|
||||
#define DISPATCH_HEAD_DIM(hd, fn, arg) \
|
||||
switch (hd) { \
|
||||
case 32: fn<32>(arg); break; \
|
||||
case 64: fn<64>(arg); break; \
|
||||
case 128: fn<128>(arg); break; \
|
||||
case 256: fn<256>(arg); break; \
|
||||
default: \
|
||||
TORCH_CHECK(false, "unsupported head_dim ", hd, \
|
||||
" (supported: 32, 64, 128, 256)"); \
|
||||
}
|
||||
|
||||
template<typename P>
|
||||
inline void alloc_split_partials(P& p) {
|
||||
auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
|
||||
auto o_part = torch::empty({p.batch, p.q_head, 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 = (float*)o_part.data_ptr();
|
||||
p.ml_part = (float*)ml_part.data_ptr();
|
||||
}
|
||||
|
||||
// ---- Shared Q-dims + strides extraction ----
|
||||
template <typename P>
|
||||
inline void extract_q_dims_and_strides(torch::Tensor& q, int64_t layout, P& p) {
|
||||
if (layout == 1) q = q.transpose(1, 2);
|
||||
p.batch = (int)q.size(0);
|
||||
p.q_head = (int)q.size(1);
|
||||
p.q_len = (int)q.size(2);
|
||||
p.head_dim = (int)q.size(3);
|
||||
p.q_stride_b = (int)q.stride(0);
|
||||
p.q_stride_h = (int)q.stride(1);
|
||||
p.q_stride_l = (int)q.stride(2);
|
||||
p.q_stride_d = (int)q.stride(3);
|
||||
}
|
||||
|
||||
// ---- Shared mask packing ----
|
||||
template <typename P>
|
||||
inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
|
||||
if (p.use_mask) {
|
||||
auto m = mask.value();
|
||||
TORCH_CHECK(m.is_cuda(), "mask must be on CUDA");
|
||||
TORCH_CHECK(m.dtype() == torch::kBool, "mask must be bool");
|
||||
TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch");
|
||||
TORCH_CHECK(m.size(m.dim() - 1) == p.kv_len, "mask kv_len mismatch");
|
||||
if (m.dim() == 2) {
|
||||
p.mask_b_stride = (int)m.stride(0);
|
||||
p.mask_q_stride = 0;
|
||||
} else if (m.dim() == 3) {
|
||||
TORCH_CHECK(m.size(1) == p.q_len, "mask q_len mismatch");
|
||||
p.mask_b_stride = (int)m.stride(0);
|
||||
p.mask_q_stride = (int)m.stride(1);
|
||||
} else {
|
||||
TORCH_CHECK(false, "mask must be 2D [batch, kv_len] or 3D [batch, q_len, kv_len]");
|
||||
}
|
||||
p.mask = m.data_ptr<bool>();
|
||||
} else {
|
||||
p.mask = nullptr;
|
||||
p.mask_b_stride = 0;
|
||||
p.mask_q_stride = 0;
|
||||
}
|
||||
}
|
||||
|
||||
// ---- attn_pack_params (contiguous KV) ----
|
||||
template<typename T>
|
||||
inline void attn_pack_params(
|
||||
torch::Tensor q,
|
||||
torch::Tensor k,
|
||||
torch::Tensor v,
|
||||
c10::optional<torch::Tensor> mask,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
int64_t layout,
|
||||
AttentionParams<T>& p
|
||||
) {
|
||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||
|
||||
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);
|
||||
TORCH_CHECK(k.sizes() == v.sizes(), "K and V must have identical shapes");
|
||||
TORCH_CHECK(q.dim() == 4 && k.dim() == 4, "Q/K/V must be 4D");
|
||||
|
||||
extract_q_dims_and_strides(q, layout, p);
|
||||
|
||||
if (layout == 1) k = k.transpose(1, 2), v = v.transpose(1, 2);
|
||||
|
||||
p.kv_head = (int)k.size(1);
|
||||
p.kv_len = (int)k.size(2);
|
||||
TORCH_CHECK(k.size(3) == p.head_dim, "K/V head_dim must match Q");
|
||||
|
||||
p.kv_stride_b = (int)k.stride(0);
|
||||
p.kv_stride_h = (int)k.stride(1);
|
||||
p.kv_stride_l = (int)k.stride(2);
|
||||
p.kv_stride_d = (int)k.stride(3);
|
||||
|
||||
p.causal_offset = (int)causal_offset;
|
||||
p.use_mask = mask.has_value() ? 1 : 0;
|
||||
p.scale = (scale > 0.0) ? (float)scale : 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();
|
||||
p.o = nullptr;
|
||||
p.o_part = nullptr;
|
||||
p.ml_part = nullptr;
|
||||
|
||||
pack_mask(mask, p);
|
||||
}
|
||||
|
||||
// ---- attn_pack_paged_params ----
|
||||
template<typename T>
|
||||
inline void attn_pack_paged_params(
|
||||
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,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
int64_t layout,
|
||||
PagedAttentionParams<T>& p
|
||||
) {
|
||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||
|
||||
TORCH_CHECK(q.is_cuda() && page_table.is_cuda() && k_cache.is_cuda() && v_cache.is_cuda());
|
||||
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
|
||||
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16");
|
||||
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16");
|
||||
TORCH_CHECK(page_table.dtype() == torch::kLong, "page_table must be int64");
|
||||
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must have identical shapes");
|
||||
|
||||
extract_q_dims_and_strides(q, layout, p);
|
||||
|
||||
p.kv_head = (int)k_cache.size(2);
|
||||
p.kv_len = (int)kv_len;
|
||||
p.page_size = (int)page_size;
|
||||
p.max_pages = (int)page_table.size(1);
|
||||
|
||||
TORCH_CHECK(q.size(2) == 1, "Q seq_len must be 1 (decode)");
|
||||
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
|
||||
TORCH_CHECK(k_cache.size(1) == page_size,
|
||||
"k_cache dim 1 must equal page_size, got ",
|
||||
k_cache.size(1), " vs ", page_size);
|
||||
|
||||
p.causal_offset = (int)causal_offset;
|
||||
p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0;
|
||||
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
|
||||
|
||||
p.page_table = page_table.data_ptr<int64_t>();
|
||||
p.k_cache = (const T*)k_cache.data_ptr();
|
||||
p.v_cache = (const T*)v_cache.data_ptr();
|
||||
p.q = (const T*)q.data_ptr();
|
||||
p.o = nullptr;
|
||||
p.o_part = nullptr;
|
||||
p.ml_part = nullptr;
|
||||
|
||||
pack_mask(mask, p);
|
||||
}
|
||||
@@ -0,0 +1,293 @@
|
||||
#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));
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Q-load: load query rows directly from global memory into mma A-operand
|
||||
// register layout. One call replaces ~15 duplicated lines in each MMA kernel.
|
||||
// stride_row is p.q_stride_h for decode (q_len=1, G heads) or
|
||||
// p.q_stride_l for prefill (multi-q rows).
|
||||
// ---------------------------------------------------------------------------
|
||||
template <int KD>
|
||||
__device__ inline void load_q_mma_frags(
|
||||
const bf16* __restrict__ q,
|
||||
int stride_row,
|
||||
int stride_d,
|
||||
int qra, int qrb,
|
||||
bool va, bool vb,
|
||||
int tid4,
|
||||
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*>(
|
||||
&q[qra * stride_row + c * stride_d]);
|
||||
const unsigned* pbu = reinterpret_cast<const unsigned*>(
|
||||
&q[qrb * stride_row + c * stride_d]);
|
||||
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;
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 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).
|
||||
// qrow0/qrow1: query row indices (for 3D mask indexing; decode passes 0).
|
||||
// mask_b_stride/mask_q_stride: mask layout (2D: mask_q_stride=0; 3D: =kv_len).
|
||||
// Reads Sacc (Q@K^T scores), applies causal/mask, computes P = exp(S - nm),
|
||||
// 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 qrow0,
|
||||
int qrow1,
|
||||
int mask_b_stride,
|
||||
int mask_q_stride,
|
||||
int mask_batch,
|
||||
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;
|
||||
int mask_base0 = mask_batch * mask_b_stride + qrow0 * mask_q_stride;
|
||||
int mask_base1 = mask_batch * mask_b_stride + qrow1 * mask_q_stride;
|
||||
#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_base0 + cc]);
|
||||
bool b1 = (c1 >= maxc0) || (has_mask && !mask[mask_base0 + c1]);
|
||||
bool b2 = (cc >= maxc1) || (has_mask && !mask[mask_base1 + cc]);
|
||||
bool b3 = (c1 >= maxc1) || (has_mask && !mask[mask_base1 + 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,82 @@
|
||||
#include "attn_paged_decode_split_kv.cuh"
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
#include "attn_paged_decode_split_kv_mma.cuh"
|
||||
#endif
|
||||
|
||||
#include "attn_entry_utils.cuh"
|
||||
|
||||
static void launch_paged_scalar_decode(PagedAttentionParams<bf16>& p) {
|
||||
int group_size = p.q_head / p.kv_head;
|
||||
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||
alloc_split_partials(p);
|
||||
|
||||
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
|
||||
dim3 grid = dim3(p.batch * p.kv_head, 1, p.num_splits);
|
||||
dim3 block = dim3(32, group_size);
|
||||
paged_attn_decode_split_kv_kernel<<<grid, block, smem>>>(p);
|
||||
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
|
||||
static void launch_paged_mma_decode(PagedAttentionParams<bf16>& p) {
|
||||
int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
||||
alloc_split_partials(p);
|
||||
|
||||
paged_attn_decode_split_kv_mma_kernel<HEAD_DIM, BC, STAGES><<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
|
||||
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
#endif
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static void dispatch_paged_decode(PagedAttentionParams<bf16>& p) {
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
int G = p.q_head / p.kv_head;
|
||||
if (G >= 1 && G <= 16 && p.page_size >= 32) {
|
||||
launch_paged_mma_decode<HEAD_DIM, 32>(p);
|
||||
return;
|
||||
}
|
||||
#endif
|
||||
launch_paged_scalar_decode(p);
|
||||
}
|
||||
|
||||
torch::Tensor attn_paged_decode(
|
||||
torch::Tensor 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,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
int64_t layout
|
||||
) {
|
||||
PagedAttentionParams<bf16> p;
|
||||
attn_pack_paged_params(q, page_table, k_cache, v_cache,
|
||||
page_size, kv_len, mask, causal_offset, scale, layout, p);
|
||||
|
||||
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
||||
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
|
||||
p.o = (bf16*)O_view.data_ptr();
|
||||
|
||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p);
|
||||
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("causal_offset") = -1,
|
||||
py::arg("scale") = 0.0,
|
||||
py::arg("layout") = 0,
|
||||
"Paged GQA decode — split-KV with direct page-table access.");
|
||||
}
|
||||
@@ -0,0 +1,147 @@
|
||||
#pragma once
|
||||
#include <cuda_bf16.h>
|
||||
#include <float.h>
|
||||
#include "attn_common.h"
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
constexpr int PDC_CHUNK = 64;
|
||||
|
||||
__device__ inline float paged_warp_reduce_sum(float val) {
|
||||
for (int offset = 16; offset > 0; offset >>= 1)
|
||||
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
|
||||
return val;
|
||||
}
|
||||
|
||||
// Split-KV scalar decode: one warp per query head, grid.z partitions KV.
|
||||
__global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p) {
|
||||
int batch = blockIdx.x / p.kv_head;
|
||||
int kv_head = blockIdx.x % p.kv_head;
|
||||
int split = blockIdx.z;
|
||||
int group_size = blockDim.y;
|
||||
int q_head = kv_head * group_size + threadIdx.y;
|
||||
int lane = threadIdx.x;
|
||||
int hd_per_thread = p.head_dim / 32;
|
||||
|
||||
// Q: stride-based [batch, q_head, q_len=1, head_dim]
|
||||
float q_reg[8];
|
||||
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
||||
+ lane * hd_per_thread * p.q_stride_d;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
q_reg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
|
||||
|
||||
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
|
||||
|
||||
extern __shared__ __align__(16) bf16 k_smem[];
|
||||
|
||||
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
||||
int chunks_per_split = (chunks_total + p.num_splits - 1) / p.num_splits;
|
||||
int ch_begin = split * chunks_per_split;
|
||||
int ch_end = min(chunks_total, ch_begin + chunks_per_split);
|
||||
|
||||
const int mask_base = batch * p.mask_b_stride;
|
||||
|
||||
for (int ci = ch_begin; ci < ch_end; ci++) {
|
||||
int chunk_start = ci * PDC_CHUNK;
|
||||
int this_chunk = min(PDC_CHUNK, p.kv_len - chunk_start);
|
||||
|
||||
int total = this_chunk * p.head_dim;
|
||||
for (int i = threadIdx.y * 32 + lane; i < total; i += blockDim.x * blockDim.y) {
|
||||
int s = i / p.head_dim;
|
||||
int d_dim = i % p.head_dim;
|
||||
int pos = chunk_start + s;
|
||||
int logical_page = pos / p.page_size;
|
||||
int page_offset = pos % p.page_size;
|
||||
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
||||
if (phys_page >= 0) {
|
||||
int64_t off = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim
|
||||
+ (int64_t)page_offset * p.kv_head * p.head_dim
|
||||
+ (int64_t)kv_head * p.head_dim
|
||||
+ d_dim;
|
||||
k_smem[i] = p.k_cache[off];
|
||||
} else {
|
||||
k_smem[i] = __float2bfloat16(0.0f);
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
for (int s = 0; s < this_chunk; s++) {
|
||||
float partial = 0.0f;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
partial += q_reg[i] * __bfloat162float(k_smem[s * p.head_dim + lane * hd_per_thread + i]);
|
||||
partial = paged_warp_reduce_sum(partial) * p.scale;
|
||||
|
||||
int kv_idx = chunk_start + s;
|
||||
if (p.use_mask && p.mask && !p.mask[mask_base + kv_idx])
|
||||
partial = -FLT_MAX;
|
||||
if (p.causal_offset >= 0 && kv_idx > p.causal_offset)
|
||||
partial = -FLT_MAX;
|
||||
|
||||
float new_m = fmaxf(m, partial);
|
||||
float alpha = expf(m - new_m);
|
||||
float beta = expf(partial - new_m);
|
||||
d = d * alpha + beta;
|
||||
|
||||
int pos = chunk_start + s;
|
||||
int logical_page = pos / p.page_size;
|
||||
int page_offset = pos % p.page_size;
|
||||
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
||||
if (phys_page >= 0) {
|
||||
int64_t v_base = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim
|
||||
+ (int64_t)page_offset * p.kv_head * p.head_dim
|
||||
+ (int64_t)kv_head * p.head_dim;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
acc_reg[i] = acc_reg[i] * alpha + __bfloat162float(p.v_cache[v_base + lane * hd_per_thread + i]) * beta;
|
||||
} else {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
acc_reg[i] = acc_reg[i] * alpha + 0.0f * beta;
|
||||
}
|
||||
m = new_m;
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
size_t bh = (size_t)batch * p.q_head + q_head;
|
||||
size_t slot = bh * p.num_splits + split;
|
||||
int d0 = lane * hd_per_thread;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
p.o_part[slot * p.head_dim + (d0 + i)] = acc_reg[i];
|
||||
if (lane == 0) {
|
||||
p.ml_part[slot * 2] = m;
|
||||
p.ml_part[slot * 2 + 1] = d;
|
||||
}
|
||||
}
|
||||
|
||||
__global__ void paged_attn_decode_combine_kernel(PagedAttentionParams<bf16> p) {
|
||||
int bh = blockIdx.x;
|
||||
int d = threadIdx.x;
|
||||
if (d >= p.head_dim) return;
|
||||
|
||||
int batch = bh / p.q_head;
|
||||
int q_head = bh % p.q_head;
|
||||
|
||||
size_t split_base = (size_t)bh * p.num_splits;
|
||||
const float* mlp = p.ml_part + split_base * 2;
|
||||
const float* op = p.o_part + split_base * p.head_dim;
|
||||
|
||||
float m = -FLT_MAX, l = 0.0f, acc = 0.0f;
|
||||
for (int s = 0; s < p.num_splits; s++) {
|
||||
float mi = mlp[s * 2];
|
||||
if (mi <= -FLT_MAX) continue;
|
||||
float li = mlp[s * 2 + 1];
|
||||
float nm = fmaxf(m, mi);
|
||||
float corr = __expf(m - nm);
|
||||
float e = __expf(mi - nm);
|
||||
acc = acc * corr + op[s * p.head_dim + d] * e;
|
||||
l = l * corr + li * e;
|
||||
m = nm;
|
||||
}
|
||||
|
||||
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
||||
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h + d * p.q_stride_d;
|
||||
p.o[o_off] = __float2bfloat16(acc * inv);
|
||||
}
|
||||
@@ -0,0 +1,170 @@
|
||||
#pragma once
|
||||
#include <cfloat>
|
||||
#include <cuda_bf16.h>
|
||||
#include "attn_common.h"
|
||||
#include "attn_mma_utils.cuh"
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
|
||||
// Paged split-KV tensor-core decode via GQA head-packing.
|
||||
// Identical algorithm to attn_decode_split_kv_mma_kernel but reads K/V
|
||||
// directly from the page pool through a page table, eliminating the gather
|
||||
// copy. Each tile (BC=32) fits within a single page (page_size >= 32), so
|
||||
// the page-table lookup happens once per tile for cp.async.
|
||||
|
||||
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
|
||||
__global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16> p) {
|
||||
constexpr int KD = HEAD_DIM / 16;
|
||||
constexpr int NC8 = BC / 8;
|
||||
constexpr int KT2 = BC / 16;
|
||||
constexpr int DN8 = HEAD_DIM / 8;
|
||||
constexpr int LD = HEAD_DIM;
|
||||
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
|
||||
constexpr int VEC = 8;
|
||||
constexpr int TOTAL = BC * HEAD_DIM;
|
||||
|
||||
const int lane = threadIdx.x;
|
||||
const int gid = lane >> 2;
|
||||
const int tid4 = lane & 3;
|
||||
|
||||
const int kv_head = blockIdx.x;
|
||||
const int batch = blockIdx.y;
|
||||
const int split = blockIdx.z;
|
||||
const int G = p.q_head / p.kv_head;
|
||||
const int q_head0 = kv_head * G;
|
||||
|
||||
__shared__ __align__(16) bf16 sK[STAGES * BC * LD];
|
||||
__shared__ __align__(16) bf16 sV[STAGES * BC * LD];
|
||||
|
||||
// ---- Load Q directly from global into mma A-operand registers ----
|
||||
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
|
||||
const int qra = gid;
|
||||
const int qrb = gid + 8;
|
||||
const bool va = qra < G, vb = qrb < G;
|
||||
unsigned Qa[KD][4];
|
||||
load_q_mma_frags<KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
|
||||
qra, qrb, va, vb, tid4, Qa);
|
||||
|
||||
float Oacc[DN8][4];
|
||||
#pragma unroll
|
||||
for (int j = 0; j < DN8; j++)
|
||||
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
||||
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
||||
|
||||
const int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
|
||||
const int ti_begin = split * tiles_per_split;
|
||||
const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
|
||||
const int has_mask = p.use_mask && p.mask;
|
||||
|
||||
// Paged strides (constant for the block)
|
||||
const int64_t page_stride = (int64_t)p.page_size * p.kv_head * HEAD_DIM;
|
||||
const int64_t pos_stride = (int64_t)p.kv_head * HEAD_DIM;
|
||||
const int64_t head_off = (int64_t)kv_head * HEAD_DIM;
|
||||
|
||||
// ---- Load tile lambda: predicated cp.async, paged addressing ----
|
||||
auto load_tile = [&](int ti, int buf) {
|
||||
int kv0 = ti * BC;
|
||||
bf16* dK = sK + buf * BC * LD;
|
||||
bf16* dV = sV + buf * BC * LD;
|
||||
int logical_page = kv0 / p.page_size;
|
||||
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
||||
bool page_valid = (phys_page >= 0);
|
||||
#pragma unroll
|
||||
for (int i = lane * VEC; i < TOTAL; i += 32 * VEC) {
|
||||
int r = i / HEAD_DIM, d = i % HEAD_DIM;
|
||||
int kc = kv0 + r;
|
||||
bool valid = (kc < p.kv_len) && page_valid;
|
||||
int page_off = kc % p.page_size;
|
||||
int64_t gmem_base = (int64_t)phys_page * page_stride
|
||||
+ (int64_t)page_off * pos_stride
|
||||
+ head_off;
|
||||
int off = r * LD + swiz_col(d, r, SWIZ_MASK);
|
||||
cp_async_16_pred(&dK[off], &p.k_cache[gmem_base + d], valid);
|
||||
cp_async_16_pred(&dV[off], &p.v_cache[gmem_base + d], valid);
|
||||
}
|
||||
cp_async_commit();
|
||||
};
|
||||
|
||||
// ---- Prologue: issue first tile load ----
|
||||
if (ti_begin < ti_end) {
|
||||
load_tile(ti_begin, 0);
|
||||
}
|
||||
|
||||
for (int ti = ti_begin; ti < ti_end; ti++) {
|
||||
constexpr int BUF_MASK = (STAGES > 1) ? (STAGES - 1) : 0;
|
||||
int buf = (ti - ti_begin) & BUF_MASK;
|
||||
|
||||
cp_async_wait_group<0>();
|
||||
__syncwarp();
|
||||
if constexpr (STAGES > 1) {
|
||||
if (ti + 1 < ti_end)
|
||||
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
|
||||
}
|
||||
|
||||
const bf16* bK = sK + buf * BC * LD;
|
||||
const bf16* bV = sV + buf * BC * LD;
|
||||
int kv0 = ti * BC;
|
||||
|
||||
float Sacc[NC8][4];
|
||||
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
|
||||
|
||||
#pragma unroll
|
||||
for (int n8 = 0; n8 < NC8; n8++)
|
||||
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||
|
||||
// Decode: q_len=1, so qrow0=qrow1=0, mask_q_stride irrelevant
|
||||
int maxc = (p.causal_offset >= 0) ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
||||
mma_softmax_tile<NC8, DN8>(kv0, maxc, maxc,
|
||||
0, 0,
|
||||
p.mask_b_stride, 0,
|
||||
batch,
|
||||
p.mask, has_mask,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
|
||||
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc);
|
||||
__syncwarp();
|
||||
|
||||
if constexpr (STAGES == 1) {
|
||||
if (ti + 1 < ti_end)
|
||||
load_tile(ti + 1, 0);
|
||||
}
|
||||
}
|
||||
|
||||
// ---- write UN-normalised partials for this split ----
|
||||
auto split_slot = [&](int h) -> size_t {
|
||||
size_t bh = (size_t)batch * p.q_head + h;
|
||||
return bh * p.num_splits + split;
|
||||
};
|
||||
#pragma unroll
|
||||
for (int dn8 = 0; dn8 < DN8; dn8++) {
|
||||
int d = dn8 * 8 + 2 * tid4;
|
||||
int r0 = gid, r1 = gid + 8;
|
||||
if (r0 < G) {
|
||||
int h = q_head0 + r0;
|
||||
float* op = p.o_part + split_slot(h) * HEAD_DIM;
|
||||
op[d] = Oacc[dn8][0];
|
||||
op[d + 1] = Oacc[dn8][1];
|
||||
}
|
||||
if (r1 < G) {
|
||||
int h = q_head0 + r1;
|
||||
float* op = p.o_part + split_slot(h) * HEAD_DIM;
|
||||
op[d] = Oacc[dn8][2];
|
||||
op[d + 1] = Oacc[dn8][3];
|
||||
}
|
||||
}
|
||||
if (tid4 == 0) {
|
||||
int r0 = gid, r1 = gid + 8;
|
||||
if (r0 < G) {
|
||||
int h = q_head0 + r0;
|
||||
float* mp = p.ml_part + split_slot(h) * 2;
|
||||
mp[0] = m0; mp[1] = l0;
|
||||
}
|
||||
if (r1 < G) {
|
||||
int h = q_head0 + r1;
|
||||
float* mp = p.ml_part + split_slot(h) * 2;
|
||||
mp[0] = m1; mp[1] = l1;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,64 @@
|
||||
#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,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
int64_t layout
|
||||
) {
|
||||
AttentionParams<bf16> p;
|
||||
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
|
||||
TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16");
|
||||
|
||||
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
||||
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
|
||||
p.o = (bf16*)O_view.data_ptr();
|
||||
|
||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_prefill, p);
|
||||
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("causal_offset") = -1,
|
||||
py::arg("scale") = 0.0,
|
||||
py::arg("layout") = 0,
|
||||
"GQA prefill (tensor-core mma on sm_80+, scalar fallback)");
|
||||
}
|
||||
@@ -0,0 +1,152 @@
|
||||
#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];
|
||||
|
||||
// Q: stride-based load [batch, q_head, q_len, head_dim]
|
||||
float qreg[DPT];
|
||||
if (q_row < p.q_len) {
|
||||
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
||||
+ q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < DPT; i++)
|
||||
qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]) * 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;
|
||||
|
||||
// KV: stride-based base
|
||||
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||
int mask_batch_base = batch * p.mask_b_stride;
|
||||
int 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);
|
||||
|
||||
// Load K/V into shared memory from strided global
|
||||
for (int i = lid; i < tlen * HEAD_DIM; i += tt) {
|
||||
int s = i / HEAD_DIM;
|
||||
int d_dim = i % HEAD_DIM;
|
||||
int kv_idx = kv0 + s;
|
||||
int g_off = kv_base + kv_idx * p.kv_stride_l + d_dim * p.kv_stride_d;
|
||||
sK[i] = p.k[g_off];
|
||||
sV[i] = p.v[g_off];
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
int lim = tlen;
|
||||
if (p.causal_offset >= 0 && 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;
|
||||
}
|
||||
|
||||
int mask_row_base = mask_batch_base + q_row * p.mask_q_stride;
|
||||
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);
|
||||
|
||||
int kv_idx = kv0 + s;
|
||||
if (p.use_mask && p.mask && !p.mask[mask_row_base + kv_idx])
|
||||
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) {
|
||||
// O: stride-based write
|
||||
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
||||
+ q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
|
||||
float rl = (l > 1e-10f) ? (1.0f / l) : 0.0f;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < DPT; i++)
|
||||
p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * rl);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,196 @@
|
||||
#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: double-buffered K/V ----
|
||||
// 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 Q fragments straight from global into mma A-operand layout.
|
||||
// stride_row = p.q_stride_l for prefill (multi-q rows across q_len).
|
||||
// See attn_mma_utils.cuh for the shared template.
|
||||
const int q_base = batch * p.q_stride_b + q_head * p.q_stride_h;
|
||||
const int 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];
|
||||
load_q_mma_frags<KD>(p.q + q_base, p.q_stride_l, p.q_stride_d,
|
||||
qra, qrb, va, vb, tid4, Qa);
|
||||
|
||||
float Oacc[DN8][4];
|
||||
#pragma unroll
|
||||
for (int j = 0; j < DN8; j++)
|
||||
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
||||
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
||||
|
||||
// KV: stride-based base
|
||||
const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||
const int tiles = (p.kv_len + BC - 1) / BC;
|
||||
const int qr0 = qrow0 + gid; // row for c0/c1
|
||||
const int qr1 = qrow0 + gid + 8; // row for c2/c3
|
||||
|
||||
// Causal tile-skip bounds (no-op when causal_offset < 0)
|
||||
const int use_skip = (p.causal_offset >= 0) ? 1 : 0;
|
||||
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;
|
||||
|
||||
// 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;
|
||||
|
||||
// ---- Load tile lambda: predicated cp.async ----
|
||||
// Issue cp.async loads for tile `ti` into shared buffer `buf`. Predicated
|
||||
// loads zero-fill rows past kv_len, so partial tiles need no scalar path.
|
||||
auto load_tile = [&](int ti, int buf) {
|
||||
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);
|
||||
int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d;
|
||||
cp_async_16_pred(&dK[off], &p.k[g_off], valid);
|
||||
cp_async_16_pred(&dV[off], &p.v[g_off], valid);
|
||||
}
|
||||
cp_async_commit();
|
||||
};
|
||||
|
||||
// ---- Prologue: issue first tile 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.causal_offset >= 0) ? min(p.kv_len, qr0 + p.causal_offset + 1)
|
||||
: p.kv_len;
|
||||
int maxc1 = (p.causal_offset >= 0) ? min(p.kv_len, qr1 + p.causal_offset + 1)
|
||||
: p.kv_len;
|
||||
mma_softmax_tile<NC8, DN8>(kv0, maxc0, maxc1,
|
||||
qr0, qr1,
|
||||
p.mask_b_stride, p.mask_q_stride,
|
||||
batch,
|
||||
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;
|
||||
// O: stride-based write
|
||||
const int o_base = batch * p.q_stride_b + q_head * p.q_stride_h;
|
||||
#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 * p.q_stride_l + d * p.q_stride_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 * p.q_stride_l + d * p.q_stride_d]) = v;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,201 @@
|
||||
/*
|
||||
Pure-C test:
|
||||
nvcc -I csrc -arch=sm_89 -O3 \
|
||||
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
|
||||
csrc/tests/attn_decode_test.cu -o test && ./test
|
||||
*/
|
||||
|
||||
#include "test_utils.cuh"
|
||||
#include "../kernels/attn_decode_split_kv.cuh"
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
#include "../kernels/attn_decode_split_kv_mma.cuh"
|
||||
#endif
|
||||
|
||||
// Split-K scratch (torch-free): the production launcher allocates these from
|
||||
// torch; here we pass pre-allocated device buffers so the bench loop doesn't
|
||||
// pay a cudaMalloc per iteration. Size for the maximum split count (32).
|
||||
struct DecodeScratch {
|
||||
float* o_part = nullptr;
|
||||
float* ml_part = nullptr;
|
||||
};
|
||||
|
||||
// Launch the production decode path (tensor-core head-packing MMA on sm_80+,
|
||||
// scalar fallback otherwise), mirroring dispatch_decode() in attn_decode.cu.
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
static bool decode_use_mma(const AttentionParams<bf16>& p) {
|
||||
int G = p.q_head / p.kv_head;
|
||||
return !p.use_mask && G > 1 && G <= 16;
|
||||
}
|
||||
|
||||
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
|
||||
static void launch_mma_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
||||
int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
||||
p.o_part = sc.o_part;
|
||||
p.ml_part = sc.ml_part;
|
||||
|
||||
attn_decode_split_kv_mma_kernel<HEAD_DIM, BC, STAGES>
|
||||
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
|
||||
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
#endif
|
||||
|
||||
static void launch_scalar_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
||||
int gs = p.q_head / p.kv_head;
|
||||
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||
p.o_part = sc.o_part;
|
||||
p.ml_part = sc.ml_part;
|
||||
|
||||
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
||||
attn_decode_split_kv_kernel<<<dim3(p.batch * p.kv_head, 1, p.num_splits), dim3(32, gs), smem>>>(p);
|
||||
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static void dispatch_decode_t(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
if (decode_use_mma(p)) { launch_mma_decode<HEAD_DIM, 32>(p, sc); return; }
|
||||
#endif
|
||||
launch_scalar_decode(p, sc);
|
||||
}
|
||||
|
||||
static void dispatch_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
||||
dispatch_by_head_dim(p.head_dim, [&]<int D>() { dispatch_decode_t<D>(p, sc); });
|
||||
}
|
||||
|
||||
// Warmed-up, CUDA-event timed sweep over the production decode MMA path.
|
||||
static void bench() {
|
||||
const int cfgs[][5] = {
|
||||
{1, 32, 4, 512, 128}, // B, Hq, Hk, kv_len, D
|
||||
{1, 32, 4, 1024, 128},
|
||||
{1, 32, 4, 2048, 128},
|
||||
{1, 32, 4, 4096, 128},
|
||||
{16, 32, 4, 2048, 128},
|
||||
{32, 32, 4, 1024, 128},
|
||||
};
|
||||
const int WARMUP = 10, ITERS = 100;
|
||||
printf("\n===== DECODE BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
|
||||
print_bench_header();
|
||||
|
||||
for (int ci = 0; ci < 6; ci++) {
|
||||
int B = cfgs[ci][0], Hq = cfgs[ci][1], Hk = cfgs[ci][2];
|
||||
int sl = cfgs[ci][3], D = cfgs[ci][4];
|
||||
size_t nQ = (size_t)B * Hq * D;
|
||||
size_t nKV = (size_t)B * Hk * sl * D;
|
||||
|
||||
bf16 *dQ, *dK, *dV, *dO;
|
||||
cudaMalloc(&dQ, nQ*2); cudaMalloc(&dK, nKV*2);
|
||||
cudaMalloc(&dV, nKV*2); cudaMalloc(&dO, nQ*2);
|
||||
size_t big = nQ > nKV ? nQ : nKV; bf16* tmp = new bf16[big];
|
||||
for (size_t i = 0; i < nQ; i++) tmp[i] = f2bf(randf());
|
||||
cudaMemcpy(dQ, tmp, nQ*2, cudaMemcpyHostToDevice);
|
||||
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
|
||||
cudaMemcpy(dK, tmp, nKV*2, cudaMemcpyHostToDevice);
|
||||
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
|
||||
cudaMemcpy(dV, tmp, nKV*2, cudaMemcpyHostToDevice);
|
||||
delete[] tmp;
|
||||
|
||||
AttentionParams<bf16> p;
|
||||
p.batch = B; p.q_head = Hq; p.kv_head = Hk; p.q_len = 1; p.kv_len = sl;
|
||||
p.head_dim = D; p.use_mask = 0; p.causal_offset = -1;
|
||||
p.scale = 1.0f / sqrtf((float)D);
|
||||
set_default_strides(p);
|
||||
p.q = dQ; p.k = dK; p.v = dV; p.mask = nullptr; p.o = dO;
|
||||
|
||||
DecodeScratch sc;
|
||||
cudaMalloc(&sc.o_part, (size_t)B*Hq*32*D*sizeof(float));
|
||||
cudaMalloc(&sc.ml_part, (size_t)B*Hq*32*2*sizeof(float));
|
||||
|
||||
auto launch = [&]() { dispatch_decode(p, sc); };
|
||||
double flops = 4.0 * B * Hq * (double)sl * D;
|
||||
double bytes = 2.0 * (2.0 * nKV * sizeof(bf16));
|
||||
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops, bytes);
|
||||
|
||||
char cfg[64];
|
||||
snprintf(cfg, sizeof(cfg),
|
||||
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
|
||||
B, Hq, Hk, 1, sl, D, 0);
|
||||
print_bench_row(cfg, r);
|
||||
|
||||
cudaFree(dQ); cudaFree(dK); cudaFree(dV); cudaFree(dO);
|
||||
cudaFree(sc.o_part); cudaFree(sc.ml_part);
|
||||
}
|
||||
}
|
||||
|
||||
int main() {
|
||||
const int configs[][5] = {
|
||||
{1, 2, 1, 64, 32}, // B,Hq,Hk,seq_len,D
|
||||
{1, 32, 4, 512, 128},
|
||||
{1, 32, 4, 1024, 128},
|
||||
};
|
||||
int n_cfgs = sizeof(configs) / sizeof(configs[0]);
|
||||
|
||||
for (int ci = 0; ci < n_cfgs; ci++) {
|
||||
int B = configs[ci][0], Hq = configs[ci][1], Hk = configs[ci][2];
|
||||
int sl = configs[ci][3], D = configs[ci][4], gs = Hq / Hk;
|
||||
printf("=== B=%d Hq=%d Hk=%d seq=%d D=%d gs=%d ===\n", B,Hq,Hk,sl,D,gs);
|
||||
|
||||
size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D;
|
||||
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
||||
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
||||
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
||||
|
||||
bool* hMask=new bool[B*sl];
|
||||
for (int i=0;i<B*sl;i++) hMask[i]=true;
|
||||
|
||||
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||
bool* dMask;
|
||||
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||
cudaMalloc(&dMask,B*sl);
|
||||
|
||||
tmp=new bf16[max(nQ,nKV)];
|
||||
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
||||
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
||||
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
||||
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
cudaMemcpy(dMask,hMask,B*sl,cudaMemcpyHostToDevice);
|
||||
|
||||
AttentionParams<bf16> p;
|
||||
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D;
|
||||
p.use_mask=0; p.causal_offset=-1;
|
||||
p.scale=1.0f/sqrtf((float)D);
|
||||
set_default_strides(p);
|
||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||
|
||||
// Split-K scratch (max 32 splits), sized for the production MMA path.
|
||||
DecodeScratch sc;
|
||||
cudaMalloc(&sc.o_part, (size_t)B*Hq*32*D*sizeof(float));
|
||||
cudaMalloc(&sc.ml_part, (size_t)B*Hq*32*2*sizeof(float));
|
||||
|
||||
double t0=now_ms();
|
||||
dispatch_decode(p, sc);
|
||||
cudaDeviceSynchronize();
|
||||
double kms=now_ms()-t0;
|
||||
cudaError_t err=cudaGetLastError();
|
||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
||||
|
||||
bf16* hOut=new bf16[nQ];
|
||||
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
||||
|
||||
float* ref=new float[nQ];
|
||||
cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, -1);
|
||||
|
||||
float max_err=0;
|
||||
for (size_t i=0;i<nQ;i++){
|
||||
float d=fabsf(bf2f(hOut[i])-ref[i]);
|
||||
if(d>max_err) max_err=d;
|
||||
}
|
||||
printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err);
|
||||
|
||||
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask);
|
||||
cudaFree(sc.o_part);cudaFree(sc.ml_part);
|
||||
delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp;
|
||||
}
|
||||
printf("All tests passed!\n");
|
||||
bench();
|
||||
return 0;
|
||||
}
|
||||
@@ -0,0 +1,332 @@
|
||||
// Compile:
|
||||
// nvcc -I csrc -arch=sm_89 -O3 --use_fast_math --ptxas-options=-O3 \
|
||||
// --extra-device-vectorization csrc/tests/attn_paged_decode_test.cu \
|
||||
// -o /tmp/test_paged && /tmp/test_paged
|
||||
|
||||
#include <cstring>
|
||||
#include "test_utils.cuh"
|
||||
#include "../kernels/attn_paged_decode_split_kv.cuh"
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
#include "../kernels/attn_paged_decode_split_kv_mma.cuh"
|
||||
#endif
|
||||
|
||||
// Copy contiguous K/V from page pool (reference gather)
|
||||
static void gather_kv_cpu(
|
||||
const bf16* h_k_pool, const bf16* h_v_pool,
|
||||
const int64_t* h_pt, int B, int Hkv, int kv_len,
|
||||
int page_size, int head_dim,
|
||||
bf16* h_k, bf16* h_v)
|
||||
{
|
||||
int max_pages = (kv_len + page_size - 1) / page_size;
|
||||
size_t page_stride = (size_t)page_size * Hkv * head_dim;
|
||||
for (int b = 0; b < B; b++) {
|
||||
for (int pos = 0; pos < kv_len; pos++) {
|
||||
int log_pg = pos / page_size;
|
||||
int pg_off = pos % page_size;
|
||||
int phys = (int)h_pt[b * max_pages + log_pg];
|
||||
for (int h = 0; h < Hkv; h++) {
|
||||
size_t src_base = (size_t)phys * page_stride
|
||||
+ (size_t)pg_off * Hkv * head_dim
|
||||
+ h * head_dim;
|
||||
size_t dst_base = ((size_t)b * Hkv + h) * kv_len * head_dim + (size_t)pos * head_dim;
|
||||
memcpy(h_k + dst_base, h_k_pool + src_base, head_dim * sizeof(bf16));
|
||||
memcpy(h_v + dst_base, h_v_pool + src_base, head_dim * sizeof(bf16));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static void launch_paged_decode(PagedAttentionParams<bf16, float>& p) {
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
int G_check = p.q_head / p.kv_head;
|
||||
bool use_mma = !p.use_mask && G_check >= 1 && G_check <= 16 && p.page_size >= 32;
|
||||
if (use_mma) {
|
||||
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
|
||||
int tiles_total = (p.kv_len + 32 - 1) / 32;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
||||
paged_attn_decode_split_kv_mma_kernel<HEAD_DIM, 32, STAGES>
|
||||
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
|
||||
} else
|
||||
#endif
|
||||
{
|
||||
int group_sz = p.q_head / p.kv_head;
|
||||
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
|
||||
paged_attn_decode_split_kv_kernel<<<
|
||||
dim3(p.batch * p.kv_head, 1, p.num_splits),
|
||||
dim3(32, group_sz), smem>>>(p);
|
||||
}
|
||||
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed) {
|
||||
printf("B=%d Hq=%d Hkv=%d kv_len=%d page_sz=%d head_dim=%d ... ", B, Hq, Hkv, kv_len, page_size, HEAD_DIM);
|
||||
fflush(stdout);
|
||||
|
||||
int max_pages = (kv_len + page_size - 1) / page_size;
|
||||
int n_phys_pages = B * max_pages;
|
||||
|
||||
size_t sz_q = (size_t)B * Hq * 1 * HEAD_DIM * sizeof(bf16);
|
||||
size_t sz_o = sz_q;
|
||||
size_t sz_kv = (size_t)n_phys_pages * page_size * Hkv * HEAD_DIM * sizeof(bf16);
|
||||
size_t sz_pt = (size_t)B * max_pages * sizeof(int64_t);
|
||||
int max_splits = 32;
|
||||
size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float);
|
||||
size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float);
|
||||
|
||||
bf16 *d_q, *d_o_paged, *d_o_ref;
|
||||
bf16 *d_k_pool, *d_v_pool;
|
||||
int64_t* d_pt;
|
||||
float *d_op, *d_ml;
|
||||
|
||||
cudaMalloc(&d_q, sz_q);
|
||||
cudaMalloc(&d_o_paged, sz_o);
|
||||
cudaMalloc(&d_o_ref, sz_o);
|
||||
cudaMalloc(&d_k_pool, sz_kv);
|
||||
cudaMalloc(&d_v_pool, sz_kv);
|
||||
cudaMalloc(&d_pt, sz_pt);
|
||||
cudaMalloc(&d_op, sz_op);
|
||||
cudaMalloc(&d_ml, sz_ml);
|
||||
|
||||
srand(seed);
|
||||
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
|
||||
|
||||
bf16* h_q = (bf16*)malloc(sz_q);
|
||||
for (int i = 0; i < B * Hq * HEAD_DIM; i++)
|
||||
h_q[i] = __float2bfloat16(rnd());
|
||||
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
|
||||
|
||||
bf16* h_k_pool = (bf16*)malloc(sz_kv);
|
||||
bf16* h_v_pool = (bf16*)malloc(sz_kv);
|
||||
size_t ps = (size_t)page_size * Hkv * HEAD_DIM;
|
||||
for (int pg = 0; pg < n_phys_pages; pg++) {
|
||||
for (int off = 0; off < page_size; off++) {
|
||||
for (int h = 0; h < Hkv; h++) {
|
||||
for (int d = 0; d < HEAD_DIM; d++) {
|
||||
float v = sinf((float)(pg * 7919 + off * 1049 + h * 331 + d));
|
||||
size_t idx = (size_t)pg * ps + (size_t)off * Hkv * HEAD_DIM + h * HEAD_DIM + d;
|
||||
h_k_pool[idx] = __float2bfloat16(v);
|
||||
h_v_pool[idx] = __float2bfloat16(v * 0.3f);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
|
||||
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
|
||||
|
||||
int64_t* h_pt = (int64_t*)malloc(sz_pt);
|
||||
int next_pg = 0;
|
||||
for (int b = 0; b < B; b++)
|
||||
for (int p = 0; p < max_pages; p++)
|
||||
h_pt[b * max_pages + p] = next_pg++;
|
||||
cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice);
|
||||
|
||||
bf16* h_k_cont = (bf16*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(bf16));
|
||||
bf16* h_v_cont = (bf16*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(bf16));
|
||||
gather_kv_cpu(h_k_pool, h_v_pool, h_pt, B, Hkv, kv_len, page_size, HEAD_DIM, h_k_cont, h_v_cont);
|
||||
|
||||
float* h_q_f = (float*)malloc((size_t)B * Hq * HEAD_DIM * sizeof(float));
|
||||
float* h_k_f = (float*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(float));
|
||||
float* h_v_f = (float*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(float));
|
||||
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
|
||||
for (int i = 0; i < B * kv_len * Hkv * HEAD_DIM; i++) {
|
||||
h_k_f[i] = bf2f(h_k_cont[i]);
|
||||
h_v_f[i] = bf2f(h_v_cont[i]);
|
||||
}
|
||||
|
||||
float* h_o_ref = (float*)calloc(B * Hq * HEAD_DIM, sizeof(float));
|
||||
cpu_attention_ref(h_q_f, h_k_f, h_v_f, nullptr, h_o_ref, B, Hq, Hkv, 1, kv_len, HEAD_DIM, -1);
|
||||
|
||||
float scale_val = 1.0f / sqrtf((float)HEAD_DIM);
|
||||
PagedAttentionParams<bf16, float> p;
|
||||
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.q_len = 1;
|
||||
p.kv_len = kv_len; p.head_dim = HEAD_DIM;
|
||||
p.use_mask = 0; p.causal_offset = -1;
|
||||
set_default_paged_strides(p);
|
||||
p.num_splits = 1; p.scale = scale_val;
|
||||
p.page_size = page_size; p.max_pages = max_pages;
|
||||
p.page_table = d_pt;
|
||||
p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||
p.q = d_q; p.mask = nullptr; p.o = d_o_paged;
|
||||
p.o_part = d_op; p.ml_part = d_ml;
|
||||
|
||||
launch_paged_decode<HEAD_DIM>(p);
|
||||
cudaDeviceSynchronize();
|
||||
|
||||
bf16* h_o_bf16 = (bf16*)malloc(sz_o);
|
||||
cudaMemcpy(h_o_bf16, d_o_paged, sz_o, cudaMemcpyDeviceToHost);
|
||||
float* h_o_paged = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
|
||||
for (int i = 0; i < B * Hq * HEAD_DIM; i++)
|
||||
h_o_paged[i] = __bfloat162float(h_o_bf16[i]);
|
||||
|
||||
float max_err = 0.0f;
|
||||
int bad_idx = -1;
|
||||
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
|
||||
float e = fabsf(h_o_paged[i] - h_o_ref[i]);
|
||||
if (e > max_err) { max_err = e; bad_idx = i; }
|
||||
}
|
||||
|
||||
bool pass = max_err < 0.02f;
|
||||
|
||||
if (pass) {
|
||||
printf("PASS (max_abs_err=%.4e)\n", max_err);
|
||||
} else {
|
||||
int b = bad_idx / (Hq * HEAD_DIM);
|
||||
int h = (bad_idx / HEAD_DIM) % Hq;
|
||||
int d = bad_idx % HEAD_DIM;
|
||||
printf("FAIL (max_abs_err=%.4e at [%d,%d,%d]: ref=%.4f got=%.4f)\n",
|
||||
max_err, b, h, d, h_o_ref[bad_idx], h_o_paged[bad_idx]);
|
||||
printf(" ref[0..7]:");
|
||||
for (int i = 0; i < 8 && i < HEAD_DIM; i++)
|
||||
printf(" %.4f", h_o_ref[i]);
|
||||
printf("\n got[0..7]:");
|
||||
for (int i = 0; i < 8 && i < HEAD_DIM; i++)
|
||||
printf(" %.4f", h_o_paged[i]);
|
||||
printf("\n");
|
||||
}
|
||||
|
||||
free(h_q); free(h_k_pool); free(h_v_pool); free(h_pt);
|
||||
free(h_k_cont); free(h_v_cont);
|
||||
free(h_q_f); free(h_k_f); free(h_v_f);
|
||||
free(h_o_ref); free(h_o_bf16); free(h_o_paged);
|
||||
cudaFree(d_q); cudaFree(d_o_paged); cudaFree(d_o_ref);
|
||||
cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_pt);
|
||||
cudaFree(d_op); cudaFree(d_ml);
|
||||
|
||||
return pass ? 0 : 1;
|
||||
}
|
||||
|
||||
struct TestCase {
|
||||
int head_dim;
|
||||
int B, Hq, Hkv, kv_len, page_size, seed;
|
||||
};
|
||||
|
||||
static const TestCase TESTS[] = {
|
||||
{128, 1, 1, 1, 8, 128, 1},
|
||||
{128, 1, 4, 4, 128, 128, 2},
|
||||
{128, 2, 4, 4, 256, 128, 3},
|
||||
{128, 1, 4, 1, 64, 64, 4},
|
||||
{128, 1, 8, 2, 64, 128, 5},
|
||||
{128, 2, 16, 4, 128, 128, 6},
|
||||
{64, 1, 4, 2, 32, 128, 7},
|
||||
{256, 1, 2, 1, 16, 128, 8},
|
||||
{32, 1, 4, 2, 32, 64, 9},
|
||||
{128, 3, 8, 2, 256, 128, 10},
|
||||
{128, 2, 32, 8, 512, 128, 11},
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
{128, 1, 16, 2, 256, 128, 12},
|
||||
{128, 2, 32, 4, 512, 128, 13},
|
||||
#endif
|
||||
};
|
||||
|
||||
static int dispatch_test(const TestCase& tc) {
|
||||
bool matched = false;
|
||||
int r = 0;
|
||||
dispatch_by_head_dim(tc.head_dim, [&]<int D>() {
|
||||
matched = true;
|
||||
r = run_test<D>(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.seed);
|
||||
});
|
||||
return matched ? r : 1;
|
||||
}
|
||||
|
||||
// Warmed-up, CUDA-event timed sweep over paged decode configs.
|
||||
// Bytes = K + V read through page table (B*Hk*kv*D each), bf16.
|
||||
template <int HEAD_DIM>
|
||||
static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) {
|
||||
int max_pages = (kv_len + page_size - 1) / page_size;
|
||||
int n_phys_pages = B * max_pages;
|
||||
|
||||
size_t sz_q = (size_t)B * Hq * 1 * HEAD_DIM * sizeof(bf16);
|
||||
size_t sz_kv = (size_t)n_phys_pages * page_size * Hkv * HEAD_DIM * sizeof(bf16);
|
||||
size_t sz_pt = (size_t)B * max_pages * sizeof(int64_t);
|
||||
int max_splits = 32;
|
||||
size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float);
|
||||
size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float);
|
||||
|
||||
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
|
||||
int64_t* d_pt;
|
||||
float *d_op, *d_ml;
|
||||
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
|
||||
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
|
||||
cudaMalloc(&d_pt, sz_pt);
|
||||
cudaMalloc(&d_op, sz_op); cudaMalloc(&d_ml, sz_ml);
|
||||
|
||||
bf16* tmp = (bf16*)malloc(sz_kv > sz_q ? sz_kv : sz_q);
|
||||
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) tmp[i] = f2bf(randf());
|
||||
cudaMemcpy(d_q, tmp, sz_q, cudaMemcpyHostToDevice);
|
||||
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) tmp[i] = f2bf(randf());
|
||||
cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
|
||||
cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
|
||||
|
||||
int64_t* h_pt = (int64_t*)malloc(sz_pt);
|
||||
int next_pg = 0;
|
||||
for (int b = 0; b < B; b++)
|
||||
for (int p = 0; p < max_pages; p++)
|
||||
h_pt[b * max_pages + p] = next_pg++;
|
||||
cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice);
|
||||
free(h_pt);
|
||||
|
||||
float scale_val = 1.0f / sqrtf((float)HEAD_DIM);
|
||||
PagedAttentionParams<bf16, float> pa;
|
||||
pa.batch = B; pa.q_head = Hq; pa.kv_head = Hkv; pa.q_len = 1;
|
||||
pa.kv_len = kv_len; pa.head_dim = HEAD_DIM;
|
||||
pa.use_mask = 0; pa.causal_offset = -1;
|
||||
set_default_paged_strides(pa);
|
||||
pa.num_splits = 1; pa.scale = scale_val;
|
||||
pa.page_size = page_size; pa.max_pages = max_pages;
|
||||
pa.page_table = d_pt;
|
||||
pa.k_cache = d_k_pool; pa.v_cache = d_v_pool;
|
||||
pa.q = d_q; pa.mask = nullptr; pa.o = d_o;
|
||||
pa.o_part = d_op; pa.ml_part = d_ml;
|
||||
|
||||
const int WARMUP = 10, ITERS = 100;
|
||||
auto launch = [&]() { launch_paged_decode<HEAD_DIM>(pa); };
|
||||
double flops = 4.0 * B * Hq * (double)kv_len * HEAD_DIM;
|
||||
size_t nKV = (size_t)B * Hkv * kv_len * HEAD_DIM;
|
||||
double bytes = 2.0 * (2.0 * nKV * sizeof(bf16));
|
||||
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops, bytes);
|
||||
|
||||
char cfg[64];
|
||||
snprintf(cfg, sizeof(cfg),
|
||||
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d page=%3d",
|
||||
B, Hq, Hkv, 1, kv_len, HEAD_DIM, page_size);
|
||||
print_bench_row(cfg, r);
|
||||
|
||||
free(tmp);
|
||||
cudaFree(d_q); cudaFree(d_o);
|
||||
cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_pt);
|
||||
cudaFree(d_op); cudaFree(d_ml);
|
||||
}
|
||||
|
||||
static void bench() {
|
||||
printf("\n===== PAGED DECODE BENCH =====\n");
|
||||
print_bench_header();
|
||||
bench_config<128>(1, 32, 4, 512, 128);
|
||||
bench_config<128>(1, 32, 4, 1024, 128);
|
||||
bench_config<128>(1, 32, 4, 2048, 128);
|
||||
bench_config<128>(1, 32, 4, 4096, 128);
|
||||
bench_config<128>(16, 32, 4, 2048, 128);
|
||||
bench_config<128>(32, 32, 4, 1024, 128);
|
||||
}
|
||||
|
||||
int main() {
|
||||
int n = sizeof(TESTS) / sizeof(TESTS[0]);
|
||||
int fail = 0;
|
||||
printf("=== Paged Decode vs CPU reference (%d cases) ===\n\n", n);
|
||||
|
||||
for (int i = 0; i < n; i++) {
|
||||
fail += dispatch_test(TESTS[i]);
|
||||
if (fail) break;
|
||||
}
|
||||
|
||||
if (fail) {
|
||||
printf("\nFAILED (%d/%d tests failed)\n", fail, n);
|
||||
return fail;
|
||||
}
|
||||
printf("\nAll %d tests passed!\n", n);
|
||||
bench();
|
||||
return 0;
|
||||
}
|
||||
@@ -0,0 +1,178 @@
|
||||
/*
|
||||
Pure-C test:
|
||||
nvcc -I csrc -arch=sm_89 -O3 \
|
||||
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
|
||||
csrc/tests/attn_prefill_test.cu -o test && ./test
|
||||
*/
|
||||
|
||||
#include "test_utils.cuh"
|
||||
#include "../kernels/attn_prefill_split_q.cuh"
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
#include "../kernels/attn_prefill_split_q_mma.cuh"
|
||||
#endif
|
||||
|
||||
// Launch the production prefill path (tensor-core MMA on sm_80+, else the
|
||||
// scalar fallback), mirroring dispatch_prefill() in attn_prefill.cu.
|
||||
template <int HEAD_DIM>
|
||||
static void launch_prefill(AttentionParams<bf16>& p) {
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
constexpr int WARPS = 4, BR = 16;
|
||||
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
|
||||
dim3 grid((p.q_len + BR * WARPS - 1) / (BR * WARPS), p.q_head, p.batch);
|
||||
dim3 block(WARPS * 32, 1, 1);
|
||||
attn_prefill_split_q_mma_kernel<HEAD_DIM, WARPS, BC><<<grid, block>>>(p);
|
||||
#else
|
||||
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
||||
dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
|
||||
dim3 block(G, ROWS, 1);
|
||||
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC><<<grid, block>>>(p);
|
||||
#endif
|
||||
}
|
||||
|
||||
static void dispatch_prefill(AttentionParams<bf16>& p) {
|
||||
switch (p.head_dim) {
|
||||
case 64: launch_prefill<64>(p); break;
|
||||
case 128: launch_prefill<128>(p); break;
|
||||
default: printf("bench: unsupported D=%d\n", p.head_dim);
|
||||
}
|
||||
}
|
||||
|
||||
// Warmed-up, CUDA-event timed throughput sweep over the production MMA path.
|
||||
// Reports per-call latency and effective tensor-core TFLOP/s (2 matmuls:
|
||||
// QK^T and P@V, each 2*B*Hq*ql*kl*D flops; halved for causal).
|
||||
static void bench() {
|
||||
const int cfgs[][7] = {
|
||||
{1,32,4,512,512,128,0},
|
||||
{1,32,4,1024,1024,128,0},
|
||||
{1,32,4,2048,2048,128,0},
|
||||
{1,32,4,2048,2048,128,1},
|
||||
{4,32,4,2048,2048,128,1},
|
||||
{1,32,4,4096,4096,128,1},
|
||||
};
|
||||
int n = sizeof(cfgs)/sizeof(cfgs[0]);
|
||||
const int WARMUP = 10, ITERS = 50;
|
||||
printf("\n===== PREFILL BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
|
||||
printf("%-46s | %10s | %10s | %10s\n",
|
||||
"config", "latency", "bandwidth", "throughput");
|
||||
printf("---------------------------------------------------------------"
|
||||
"----------------------------\n");
|
||||
|
||||
for (int ci = 0; ci < n; ci++) {
|
||||
int B=cfgs[ci][0], Hq=cfgs[ci][1], Hk=cfgs[ci][2];
|
||||
int ql=cfgs[ci][3], kl=cfgs[ci][4], D=cfgs[ci][5], causal=cfgs[ci][6];
|
||||
size_t nQ=(size_t)B*Hq*ql*D, nKV=(size_t)B*Hk*kl*D;
|
||||
|
||||
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||
size_t big = nQ>nKV?nQ:nKV; tmp=new bf16[big];
|
||||
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(randf());
|
||||
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
|
||||
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
|
||||
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
|
||||
AttentionParams<bf16> p;
|
||||
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
|
||||
p.use_mask=0; p.causal_offset=causal?0:-1;
|
||||
set_default_strides(p);
|
||||
p.scale=1.0f/sqrtf((float)D);
|
||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||
|
||||
for (int i=0;i<WARMUP;i++) dispatch_prefill(p);
|
||||
cudaDeviceSynchronize();
|
||||
cudaError_t err=cudaGetLastError();
|
||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return;}
|
||||
|
||||
cudaEvent_t s,e; cudaEventCreate(&s); cudaEventCreate(&e);
|
||||
cudaEventRecord(s);
|
||||
for (int i=0;i<ITERS;i++) dispatch_prefill(p);
|
||||
cudaEventRecord(e); cudaEventSynchronize(e);
|
||||
float ms=0; cudaEventElapsedTime(&ms,s,e); ms/=ITERS;
|
||||
|
||||
double flops = 4.0*B*Hq*(double)ql*kl*D;
|
||||
if (causal) flops *= 0.5;
|
||||
double tflops = flops/(ms*1e-3)/1e12;
|
||||
// HBM traffic: Q + O (B*Hq*ql*D each) + K + V (B*Hk*kl*D each), bf16.
|
||||
double bytes = 2.0 * (2.0*nQ + 2.0*nKV);
|
||||
double gbps = bytes/(ms*1e-3)/1e9;
|
||||
|
||||
char cfg[64];
|
||||
snprintf(cfg, sizeof(cfg),
|
||||
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
|
||||
B,Hq,Hk,ql,kl,D,causal);
|
||||
printf("%-46s | %7.4f ms | %7.1f GB/s | %6.2f TFLOP/s\n",
|
||||
cfg, ms, gbps, tflops);
|
||||
|
||||
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
|
||||
delete[]tmp; cudaEventDestroy(s); cudaEventDestroy(e);
|
||||
}
|
||||
}
|
||||
|
||||
int main() {
|
||||
const int configs[][7] = {
|
||||
{1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal
|
||||
{1,32,4,512,512,128,0}, // standard
|
||||
{1,32,4,128,256,128,0}, // medium
|
||||
{1,4,2,256,256,128,1}, // causal
|
||||
};
|
||||
int n_configs = sizeof(configs) / sizeof(configs[0]);
|
||||
|
||||
for (int ci = 0; ci < n_configs; ci++) {
|
||||
int B=configs[ci][0], Hq=configs[ci][1], Hk=configs[ci][2];
|
||||
int ql=configs[ci][3], kl=configs[ci][4], D=configs[ci][5];
|
||||
int causal=configs[ci][6];
|
||||
printf("=== B=%d Hq=%d Hk=%d q=%d kv=%d D=%d causal=%d ===\n",
|
||||
B,Hq,Hk,ql,kl,D,causal);
|
||||
|
||||
size_t nQ = B*Hq*ql*D, nKV = B*Hk*kl*D;
|
||||
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
||||
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
||||
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
||||
|
||||
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||
tmp=new bf16[max(nQ,nKV)];
|
||||
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
||||
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
||||
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
||||
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
|
||||
AttentionParams<bf16> p;
|
||||
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
|
||||
p.use_mask=0; p.causal_offset=causal?0:-1;
|
||||
set_default_strides(p);
|
||||
p.scale=1.0f/sqrtf((float)D);
|
||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||
|
||||
double t0=now_ms();
|
||||
dispatch_prefill(p);
|
||||
cudaDeviceSynchronize();
|
||||
double kms=now_ms()-t0;
|
||||
cudaError_t err=cudaGetLastError();
|
||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
||||
|
||||
bf16* hOut=new bf16[nQ];
|
||||
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
||||
|
||||
float* ref=new float[nQ];
|
||||
cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal ? 0 : -1);
|
||||
|
||||
float max_err=0;
|
||||
for (size_t i=0;i<nQ;i++) {
|
||||
float d=fabsf(bf2f(hOut[i])-ref[i]);
|
||||
if(d>max_err) max_err=d;
|
||||
}
|
||||
printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err);
|
||||
|
||||
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
|
||||
delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp;
|
||||
}
|
||||
printf("All tests passed!\n");
|
||||
bench();
|
||||
return 0;
|
||||
}
|
||||
@@ -0,0 +1,181 @@
|
||||
#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);
|
||||
}
|
||||
|
||||
// Set default strides for contiguous b h l d layout on AttentionParams.
|
||||
template<typename P>
|
||||
inline void set_default_strides(P& p) {
|
||||
p.q_stride_b = p.q_head * p.q_len * p.head_dim;
|
||||
p.q_stride_h = p.q_len * p.head_dim;
|
||||
p.q_stride_l = p.head_dim;
|
||||
p.q_stride_d = 1;
|
||||
p.kv_stride_b = p.kv_head * p.kv_len * p.head_dim;
|
||||
p.kv_stride_h = p.kv_len * p.head_dim;
|
||||
p.kv_stride_l = p.head_dim;
|
||||
p.kv_stride_d = 1;
|
||||
p.mask_b_stride = p.kv_len;
|
||||
p.mask_q_stride = 0;
|
||||
}
|
||||
|
||||
// Set default Q strides for contiguous b h l d layout on PagedAttentionParams.
|
||||
template<typename P>
|
||||
inline void set_default_paged_strides(P& p) {
|
||||
p.q_stride_b = p.q_head * p.q_len * p.head_dim;
|
||||
p.q_stride_h = p.q_len * p.head_dim;
|
||||
p.q_stride_l = p.head_dim;
|
||||
p.q_stride_d = 1;
|
||||
p.mask_b_stride = p.kv_len;
|
||||
p.mask_q_stride = 0;
|
||||
}
|
||||
|
||||
// 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.
|
||||
// causal_offset: -1 = non-causal; >=0 = absolute position of first Q token.
|
||||
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 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 (causal_offset >= 0) {
|
||||
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;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
+3
-3
@@ -9,8 +9,8 @@ readme = "README.md"
|
||||
requires-python = ">=3.12"
|
||||
dependencies = [
|
||||
"h5py==3.15.1",
|
||||
"numpy==2.3.2",
|
||||
"torch==2.7.1",
|
||||
"numpy==2.4.4",
|
||||
"torch==2.11.0",
|
||||
"tokenizers==0.21.4",
|
||||
"tqdm==4.67.1",
|
||||
"safetensors==0.5.3",
|
||||
@@ -37,7 +37,7 @@ dev = ["pytest==9.0.2", "ruff"]
|
||||
where = ["."]
|
||||
|
||||
[tool.pip]
|
||||
extra-index-url = "https://download.pytorch.org/whl/cu126"
|
||||
extra-index-url = "https://download.pytorch.org/whl/cu128"
|
||||
|
||||
[tool.setuptools.dynamic]
|
||||
version = { attr = "astrai.__version__" }
|
||||
|
||||
@@ -5,7 +5,7 @@ from huggingface_hub import snapshot_download
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[2]
|
||||
DEFAULT_LOCAL_DIR = Path(PROJECT_ROOT, "params")
|
||||
DEFAULT_REPO_ID = "ViperEk/KHAOSZ"
|
||||
DEFAULT_REPO_ID = "ViperEkura/AstrAI-V1-instruct"
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(
|
||||
|
||||
@@ -26,11 +26,9 @@ def batch_generate():
|
||||
|
||||
prompts = [
|
||||
tokenizer.apply_chat_template(
|
||||
[
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": q},
|
||||
],
|
||||
[{"role": "user", "content": q}],
|
||||
tokenize=False,
|
||||
add_generation_prompt=True,
|
||||
)
|
||||
for q in inputs
|
||||
]
|
||||
|
||||
+73
-16
@@ -1,3 +1,4 @@
|
||||
from argparse import ArgumentParser
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
@@ -7,15 +8,69 @@ from astrai.model import AutoModel
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[2]
|
||||
PARAMETER_ROOT = Path(PROJECT_ROOT, "params")
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = ArgumentParser(description="Interactive streaming chat")
|
||||
parser.add_argument(
|
||||
"--model_path",
|
||||
type=Path,
|
||||
default=PROJECT_ROOT / "params",
|
||||
help="Path to model weights (params/ or checkpoint/epoch_N_step_M/)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--temperature",
|
||||
type=float,
|
||||
default=0.8,
|
||||
help="Sampling temperature (default: 0.8)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--top_p",
|
||||
type=float,
|
||||
default=0.95,
|
||||
help="Top-p sampling threshold",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--top_k",
|
||||
type=int,
|
||||
default=50,
|
||||
help="Top-k sampling threshold",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_tokens",
|
||||
type=int,
|
||||
default=2048,
|
||||
help="Maximum tokens to generate",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--frequency_penalty",
|
||||
type=float,
|
||||
default=0.5,
|
||||
help="Penalty per occurrence for repeated tokens (0.0 disables, "
|
||||
"range -2.0~2.0, typical 0.3-1.0)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--rep_window",
|
||||
type=int,
|
||||
default=64,
|
||||
help="Number of recent prompt tokens to include in penalty history",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--system_prompt",
|
||||
type=str,
|
||||
default="",
|
||||
help="Optional system prompt (default: empty, model not SFT-trained on system role)",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def chat():
|
||||
model = AutoModel.from_pretrained(PARAMETER_ROOT)
|
||||
tokenizer = AutoTokenizer.from_pretrained(PARAMETER_ROOT)
|
||||
model.to(device="cuda", dtype=torch.bfloat16)
|
||||
args = parse_args()
|
||||
model_path = args.model_path
|
||||
|
||||
messages = [{"role": "system", "content": "You are a helpful assistant."}]
|
||||
model = AutoModel.from_pretrained(model_path)
|
||||
tokenizer = AutoTokenizer.from_pretrained(model_path)
|
||||
model.to(device="cuda", dtype=torch.bfloat16)
|
||||
engine = InferenceEngine(model=model, tokenizer=tokenizer)
|
||||
|
||||
while True:
|
||||
@@ -23,27 +78,29 @@ def chat():
|
||||
if query == "!exit":
|
||||
break
|
||||
|
||||
# Add user message
|
||||
messages.append({"role": "user", "content": query})
|
||||
msgs = []
|
||||
if args.system_prompt:
|
||||
msgs.append({"role": "system", "content": args.system_prompt})
|
||||
msgs.append({"role": "user", "content": query})
|
||||
prompt = tokenizer.apply_chat_template(
|
||||
msgs, tokenize=False, add_generation_prompt=True
|
||||
)
|
||||
|
||||
# Generate response
|
||||
full_response = ""
|
||||
prompt = tokenizer.apply_chat_template(messages, tokenize=False)
|
||||
|
||||
for token in engine.generate(
|
||||
prompt=prompt,
|
||||
stream=True,
|
||||
max_tokens=2048,
|
||||
temperature=0.8,
|
||||
top_p=0.95,
|
||||
top_k=50,
|
||||
max_tokens=args.max_tokens,
|
||||
temperature=args.temperature,
|
||||
top_p=args.top_p,
|
||||
top_k=args.top_k,
|
||||
frequency_penalty=args.frequency_penalty,
|
||||
rep_window=args.rep_window,
|
||||
):
|
||||
print(token, end="", flush=True)
|
||||
full_response += token
|
||||
|
||||
print()
|
||||
# Add assistant response to messages
|
||||
messages.append({"role": "assistant", "content": full_response.strip()})
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -0,0 +1,321 @@
|
||||
"""SVD effective rank & weight statistics analysis for model checkpoints."""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import safetensors.torch
|
||||
import torch
|
||||
|
||||
|
||||
def effective_rank_metrics(w: torch.Tensor) -> dict:
|
||||
if w.ndim == 1:
|
||||
return {"shape": tuple(w.shape), "is_1d": True}
|
||||
|
||||
w = w.float()
|
||||
s = torch.linalg.svdvals(w)
|
||||
s_sq = s**2
|
||||
total = s_sq.sum()
|
||||
cumsum = torch.cumsum(s_sq, dim=0) / total
|
||||
|
||||
min_dim = min(w.shape[0], w.shape[1])
|
||||
er_90 = (cumsum < 0.90).sum().item() + 1
|
||||
er_95 = (cumsum < 0.95).sum().item() + 1
|
||||
er_99 = (cumsum < 0.99).sum().item() + 1
|
||||
|
||||
p = s_sq / total
|
||||
p = p[p > 1e-30]
|
||||
entropy = -(p * torch.log(p)).sum()
|
||||
entropic_rank = torch.exp(entropy).item()
|
||||
|
||||
return {
|
||||
"shape": tuple(w.shape),
|
||||
"min_dim": min_dim,
|
||||
"er_90": er_90,
|
||||
"er_95": er_95,
|
||||
"er_99": er_99,
|
||||
"er_99_norm": er_99 / min_dim,
|
||||
"er_95_norm": er_95 / min_dim,
|
||||
"entropic_rank": entropic_rank,
|
||||
"entropic_rank_norm": entropic_rank / min_dim,
|
||||
"top1_ratio": s[0].item() / s.sum().item(),
|
||||
"top5_ratio": s[:5].sum().item() / s.sum().item(),
|
||||
"decay_ratio": s[-1].item() / s[0].item(),
|
||||
"condition_number": s[0].item() / s[-1].item(),
|
||||
"mean": w.mean().item(),
|
||||
"std": w.std().item(),
|
||||
"min": w.min().item(),
|
||||
"max": w.max().item(),
|
||||
}
|
||||
|
||||
|
||||
def format_header(headers: list[str], widths: list[int]) -> str:
|
||||
return "".join(h.ljust(w) for h, w in zip(headers, widths))
|
||||
|
||||
|
||||
def format_row(values: list[str], widths: list[int]) -> str:
|
||||
return "".join(v.ljust(w) for v, w in zip(values, widths))
|
||||
|
||||
|
||||
def group_by_component(results: dict[str, dict]) -> dict[str, list[dict]]:
|
||||
groups: dict[str, list[dict]] = {}
|
||||
for key, r in results.items():
|
||||
parts = key.split(".")
|
||||
if parts[0] == "layers" and len(parts) >= 3:
|
||||
sub = parts[2:]
|
||||
if sub[0] == "attention":
|
||||
comp = f"attn.{sub[1]}"
|
||||
elif sub[0] == "mlp":
|
||||
comp = f"mlp.{sub[1]}"
|
||||
elif sub[0] == "input_norm":
|
||||
comp = "input_norm"
|
||||
elif sub[0] == "post_attention_norm":
|
||||
comp = "post_attn_norm"
|
||||
else:
|
||||
comp = ".".join(sub)
|
||||
else:
|
||||
comp = key
|
||||
groups.setdefault(comp, []).append(r)
|
||||
return groups
|
||||
|
||||
|
||||
def print_component_summary(results: dict[str, dict], title: str):
|
||||
groups = group_by_component(results)
|
||||
matrix_groups = {
|
||||
k: [v for v in vs if not v.get("is_1d")]
|
||||
for k, vs in groups.items()
|
||||
if any(not v.get("is_1d") for v in vs)
|
||||
}
|
||||
|
||||
widths = [20, 12, 12, 12, 12, 12]
|
||||
print(f"\n{title}")
|
||||
print(
|
||||
format_header(
|
||||
["Component", "N", "ER@99%", "EntRank%", "Top1 σ(%)", "Cond. Num"], widths
|
||||
)
|
||||
)
|
||||
print("-" * sum(widths))
|
||||
|
||||
for name in sorted(matrix_groups.keys()):
|
||||
items = matrix_groups[name]
|
||||
n = len(items)
|
||||
print(
|
||||
format_row(
|
||||
[
|
||||
name,
|
||||
str(n),
|
||||
f"{sum(r['er_99_norm'] for r in items) / n:.4f}",
|
||||
f"{sum(r['entropic_rank_norm'] for r in items) / n:.4f}",
|
||||
f"{sum(r['top1_ratio'] for r in items) / n:.4f}",
|
||||
f"{sum(r['condition_number'] for r in items) / n:.1f}",
|
||||
],
|
||||
widths,
|
||||
)
|
||||
)
|
||||
|
||||
all_er = [
|
||||
r["er_99_norm"]
|
||||
for vs in matrix_groups.values()
|
||||
for r in vs
|
||||
if not r.get("is_1d")
|
||||
]
|
||||
if all_er:
|
||||
m = sum(all_er) / len(all_er)
|
||||
print(f"\n Overall Mean ER@99: {m:.4f} ({m * 100:.1f}% of dimension)")
|
||||
if m > 0.85:
|
||||
print(" → HIGH utilization: model near capacity → need more params")
|
||||
elif m > 0.5:
|
||||
print(" → MODERATE utilization: some headroom left")
|
||||
else:
|
||||
print(" → LOW utilization: significant unused capacity")
|
||||
|
||||
|
||||
def print_layer_grid(results: dict[str, dict]):
|
||||
comps = [
|
||||
"attn.q_proj",
|
||||
"attn.k_proj",
|
||||
"attn.v_proj",
|
||||
"attn.o_proj",
|
||||
"mlp.up",
|
||||
"mlp.gate",
|
||||
"mlp.down",
|
||||
]
|
||||
widths = [6] + [10] * len(comps)
|
||||
metric = "er_99_norm"
|
||||
|
||||
print(f"\n--- Per-Layer Effective Rank (99% energy) ---")
|
||||
print(format_header(["Layer"] + comps, widths))
|
||||
print("-" * sum(widths))
|
||||
|
||||
layer_data: dict[int, dict[str, dict]] = {}
|
||||
for key, r in results.items():
|
||||
parts = key.split(".")
|
||||
if parts[0] != "layers":
|
||||
continue
|
||||
li = int(parts[1])
|
||||
sub = parts[2:]
|
||||
if sub[0] == "attention":
|
||||
cname = f"attn.{sub[1]}"
|
||||
elif sub[0] == "mlp":
|
||||
cname = f"mlp.{sub[1]}"
|
||||
else:
|
||||
continue
|
||||
layer_data.setdefault(li, {})[cname] = r
|
||||
|
||||
for li in sorted(layer_data):
|
||||
values = [str(li)]
|
||||
for c in comps:
|
||||
v = layer_data[li].get(c, {}).get(metric, 0)
|
||||
values.append(f"{v:.4f}")
|
||||
print(format_row(values, widths))
|
||||
|
||||
|
||||
def print_weight_stats(results: dict[str, dict]):
|
||||
groups = group_by_component(results)
|
||||
widths = [20, 12, 12, 12, 12]
|
||||
print(f"\n--- Weight Value Statistics ---")
|
||||
print(format_header(["Component", "Mean", "Std", "Min", "Max"], widths))
|
||||
print("-" * sum(widths))
|
||||
|
||||
for name in sorted(groups.keys()):
|
||||
items = groups[name]
|
||||
means = [r.get("mean", 0) for r in items]
|
||||
stds = [r.get("std", 0) for r in items]
|
||||
mins = [r.get("min", 0) for r in items]
|
||||
maxs = [r.get("max", 0) for r in items]
|
||||
g_mean = sum(means) / len(means)
|
||||
g_std = sum(stds) / len(stds)
|
||||
g_min = min(mins)
|
||||
g_max = max(maxs)
|
||||
print(
|
||||
format_row(
|
||||
[
|
||||
name,
|
||||
f"{g_mean:.6f}",
|
||||
f"{g_std:.6f}",
|
||||
f"{g_min:.6f}",
|
||||
f"{g_max:.6f}",
|
||||
],
|
||||
widths,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def print_params_summary(results: dict[str, dict]):
|
||||
total_2d = sum(
|
||||
r["shape"][0] * r["shape"][1] for r in results.values() if not r.get("is_1d")
|
||||
)
|
||||
total_1d = sum(r["shape"][0] for r in results.values() if r.get("is_1d"))
|
||||
print(f"\n Total 2D params: {total_2d:,}")
|
||||
print(f" Total 1D params: {total_1d:,}")
|
||||
print(f" Total params: {total_2d + total_1d:,}")
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="SVD effective rank & weight statistics of a model checkpoint."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ckpt_dir",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to checkpoint directory (containing model.safetensors + config.json).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--compare",
|
||||
type=str,
|
||||
nargs="*",
|
||||
help="Additional checkpoint directories to compare against.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--no_svd",
|
||||
action="store_true",
|
||||
help="Skip SVD analysis, only show weight statistics (mean/std/min/max).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Save results as JSON to this path.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
all_results = {}
|
||||
|
||||
def analyze_one(ckpt_dir: str, label: str):
|
||||
ckpt_dir = Path(ckpt_dir)
|
||||
weights_path = ckpt_dir / "model.safetensors"
|
||||
if not weights_path.exists():
|
||||
print(f"ERROR: {weights_path} not found")
|
||||
return {}
|
||||
|
||||
meta = {}
|
||||
meta_path = ckpt_dir / "meta.json"
|
||||
if meta_path.exists():
|
||||
with open(meta_path) as f:
|
||||
meta = json.load(f)
|
||||
|
||||
print(f"\n{'=' * 70}")
|
||||
print(f" {label}: {ckpt_dir}")
|
||||
if meta:
|
||||
print(
|
||||
f" Iteration: {meta.get('iteration', '?')}, "
|
||||
f"Strategy: {meta.get('strategy', '?')}, "
|
||||
f"nprocs={meta.get('nprocs', '?')}"
|
||||
)
|
||||
print(f"{'=' * 70}")
|
||||
|
||||
print(f"Loading weights...")
|
||||
sd = safetensors.torch.load_file(str(weights_path))
|
||||
print(f" {len(sd)} keys loaded")
|
||||
|
||||
weight_keys = [
|
||||
k
|
||||
for k in sd
|
||||
if ".weight" in k and "rotary_embedding" not in k and "freqs_cis" not in k
|
||||
]
|
||||
|
||||
results = {}
|
||||
if not args.no_svd:
|
||||
print(f"Computing SVD on {len(weight_keys)} tensors...")
|
||||
for i, k in enumerate(sorted(weight_keys)):
|
||||
print(f" [{i + 1}/{len(weight_keys)}] {k:<60s}", end="\r")
|
||||
results[k] = effective_rank_metrics(sd[k])
|
||||
print()
|
||||
else:
|
||||
print(f"Computing stats on {len(weight_keys)} tensors (no SVD)...")
|
||||
for i, k in enumerate(sorted(weight_keys)):
|
||||
t = sd[k]
|
||||
results[k] = {
|
||||
"shape": tuple(t.shape),
|
||||
"is_1d": t.ndim == 1,
|
||||
"mean": t.float().mean().item(),
|
||||
"std": t.float().std().item(),
|
||||
"min": t.float().min().item(),
|
||||
"max": t.float().max().item(),
|
||||
}
|
||||
|
||||
print_params_summary(results)
|
||||
if not args.no_svd:
|
||||
print_component_summary(
|
||||
results, "\n=== SVD Effective Rank by Component ==="
|
||||
)
|
||||
print_layer_grid(results)
|
||||
print_weight_stats(results)
|
||||
all_results[label] = results
|
||||
return results
|
||||
|
||||
analyze_one(args.ckpt_dir, "Primary")
|
||||
|
||||
if args.compare:
|
||||
for cdir in args.compare:
|
||||
analyze_one(cdir, f"Compare_{cdir}")
|
||||
|
||||
if args.output:
|
||||
with open(args.output, "w", encoding="utf-8") as f:
|
||||
json.dump(all_results, f, indent=2)
|
||||
print(f"\nResults saved to {args.output}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
+295
-223
@@ -1,36 +1,38 @@
|
||||
"""HumanEval code generation benchmark.
|
||||
"""HumanEval benchmark — functional pipeline design.
|
||||
|
||||
Generates n completions per problem, extracts function bodies, executes
|
||||
against hidden tests, and computes pass@k.
|
||||
Pipeline:
|
||||
load -> generate -> extract -> test -> score -> report
|
||||
|
||||
Usage::
|
||||
|
||||
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
|
||||
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.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import subprocess
|
||||
import sys
|
||||
from dataclasses import dataclass
|
||||
from math import prod
|
||||
from multiprocessing import Process, Queue
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
from typing import Dict, Iterator, List, Optional, Sequence, Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import tqdm
|
||||
from datasets import load_dataset
|
||||
|
||||
from astrai.inference import InferenceEngine
|
||||
from astrai.model import AutoModel
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
HUMANEVAL_URL = (
|
||||
"https://github.com/openai/human-eval/raw/master/data/HumanEval.jsonl.gz"
|
||||
)
|
||||
# ---------------------------------------------------------------------------
|
||||
# Config
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_STOP_SEQUENCES = [
|
||||
HUMANEVAL_HF_DATASET = "openai/openai_humaneval"
|
||||
|
||||
STOP_SEQUENCES = [
|
||||
"\nclass ",
|
||||
"\ndef ",
|
||||
"\n# ",
|
||||
@@ -40,43 +42,80 @@ _STOP_SEQUENCES = [
|
||||
]
|
||||
|
||||
|
||||
def _download_humaneval(data_path: str):
|
||||
if os.path.exists(data_path):
|
||||
@dataclass
|
||||
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(path: str):
|
||||
if os.path.exists(path):
|
||||
return
|
||||
import gzip
|
||||
import urllib.request
|
||||
|
||||
os.makedirs(os.path.dirname(data_path) or ".", exist_ok=True)
|
||||
print(f"Downloading HumanEval from {HUMANEVAL_URL} ...")
|
||||
tmp = data_path + ".tmp"
|
||||
urllib.request.urlretrieve(HUMANEVAL_URL, tmp)
|
||||
with gzip.open(tmp, "rb") as f_in:
|
||||
with open(data_path, "wb") as f_out:
|
||||
f_out.write(f_in.read())
|
||||
os.remove(tmp)
|
||||
print(f" saved to {data_path}")
|
||||
os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
|
||||
print(f"Downloading HumanEval from HuggingFace ({HUMANEVAL_HF_DATASET}) ...")
|
||||
ds = load_dataset(HUMANEVAL_HF_DATASET, split="test")
|
||||
with open(path, "w", encoding="utf-8") as f:
|
||||
for item in ds:
|
||||
f.write(json.dumps(item, ensure_ascii=False) + "\n")
|
||||
print(f" saved {len(ds)} problems to {path}")
|
||||
|
||||
|
||||
def _load_problems(data_path: str) -> List[dict]:
|
||||
problems = []
|
||||
with open(data_path, "r", encoding="utf-8") as f:
|
||||
def load_jsonl(path: str) -> List[dict]:
|
||||
rows = []
|
||||
with open(path, encoding="utf-8") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if line:
|
||||
problems.append(json.loads(line))
|
||||
return problems
|
||||
rows.append(json.loads(line))
|
||||
return rows
|
||||
|
||||
|
||||
def _extract_function_body(code: str, entry_point: str) -> Optional[str]:
|
||||
"""Extract the function body from a completion."""
|
||||
def save_json(path: str, data):
|
||||
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[^:]*:"
|
||||
match = re.search(pattern, code)
|
||||
if not match:
|
||||
# Use the full code as-is if we can't find the function
|
||||
return code
|
||||
|
||||
body_start = match.end()
|
||||
lines = code[body_start:].split("\n")
|
||||
lines = code[match.end() :].split("\n")
|
||||
body_lines = []
|
||||
started = False
|
||||
|
||||
@@ -94,240 +133,273 @@ def _extract_function_body(code: str, entry_point: str) -> Optional[str]:
|
||||
body_lines.append(stripped)
|
||||
|
||||
body = "\n".join(body_lines)
|
||||
if not body.strip():
|
||||
return None
|
||||
return body
|
||||
return body if body.strip() else None
|
||||
|
||||
|
||||
def _trim_stop_sequences(text: str) -> 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]:
|
||||
def deduplicate(seq: Sequence[str]) -> List[str]:
|
||||
seen = set()
|
||||
unique = []
|
||||
for c in completions:
|
||||
if c not in seen:
|
||||
seen.add(c)
|
||||
unique.append(c)
|
||||
return unique
|
||||
return [x for x in seq if not (x in seen or seen.add(x))]
|
||||
|
||||
|
||||
def _generate(
|
||||
def generate_batch(
|
||||
engine: InferenceEngine,
|
||||
prompt: str,
|
||||
num_samples: int,
|
||||
n: int,
|
||||
batch_size: int,
|
||||
max_tokens: int,
|
||||
temperature: float,
|
||||
top_p: float,
|
||||
top_k: int,
|
||||
batch_size: int,
|
||||
) -> List[str]:
|
||||
batches = [prompt] * min(batch_size, num_samples)
|
||||
completions = []
|
||||
remaining = num_samples
|
||||
|
||||
remaining = n
|
||||
while remaining > 0:
|
||||
current = min(batch_size, remaining)
|
||||
batch_prompts = batches[:current]
|
||||
outputs = engine.generate(
|
||||
prompt=batch_prompts,
|
||||
prompt=[prompt] * current,
|
||||
stream=False,
|
||||
max_tokens=max_tokens,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
)
|
||||
if isinstance(outputs, str):
|
||||
outputs = [outputs]
|
||||
completions.extend(outputs)
|
||||
completions.extend(outputs if isinstance(outputs, list) else [outputs])
|
||||
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,
|
||||
problems: List[dict],
|
||||
num_samples: int,
|
||||
max_tokens: int,
|
||||
temperature: float,
|
||||
top_p: float,
|
||||
top_k: int,
|
||||
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(
|
||||
problems: Sequence[dict],
|
||||
cfg: EvalConfig,
|
||||
) -> List[dict]:
|
||||
results = []
|
||||
for problem in tqdm.tqdm(problems, desc="Generating", unit="problem"):
|
||||
raw = generate_batch(
|
||||
engine,
|
||||
prompt,
|
||||
num_samples,
|
||||
max_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
batch_size,
|
||||
problem["prompt"],
|
||||
cfg.num_samples,
|
||||
cfg.batch_size,
|
||||
cfg.max_tokens,
|
||||
cfg.temperature,
|
||||
cfg.top_p,
|
||||
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
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="HumanEval benchmark")
|
||||
parser.add_argument(
|
||||
"--param_path", type=str, default="./params", help="Model directory"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--data_path",
|
||||
def execute_one(args: tuple) -> bool:
|
||||
full_code, entry_point, timeout = args
|
||||
try:
|
||||
r = subprocess.run(
|
||||
[sys.executable, "-c", full_code],
|
||||
capture_output=True,
|
||||
timeout=timeout,
|
||||
)
|
||||
return r.returncode == 0
|
||||
except subprocess.TimeoutExpired:
|
||||
return False
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def test_one(item: dict, cfg: EvalConfig, pool=None) -> 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)
|
||||
|
||||
def _run(p):
|
||||
return sum(1 for ok in p.map(execute_one, codes) if ok)
|
||||
|
||||
if pool is not None:
|
||||
passed = _run(pool)
|
||||
else:
|
||||
with ProcessPoolExecutor(max_workers=cfg.test_workers) as p:
|
||||
passed = _run(p)
|
||||
|
||||
return task_id, n, passed
|
||||
|
||||
|
||||
def test_all(
|
||||
items: Sequence[dict],
|
||||
cfg: EvalConfig,
|
||||
) -> Iterator[Tuple[str, int, int]]:
|
||||
from concurrent.futures import ProcessPoolExecutor
|
||||
|
||||
pool = ProcessPoolExecutor(max_workers=cfg.test_workers)
|
||||
try:
|
||||
for item in tqdm.tqdm(items, desc="Testing", unit="problem"):
|
||||
yield test_one(item, cfg, pool)
|
||||
finally:
|
||||
pool.shutdown(wait=True)
|
||||
|
||||
|
||||
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:
|
||||
"""Score pass@k for each problem.
|
||||
|
||||
k values are filtered per-problem: if a problem has n < k samples
|
||||
(e.g. after deduplication), pass@k is not computed for that problem.
|
||||
The summary averages only over problems where the k was computed.
|
||||
"""
|
||||
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:
|
||||
if k <= n:
|
||||
pk = round(pass_at_k(n, passed, k), 4)
|
||||
entry[f"pass@{k}"] = pk
|
||||
scores[k].append(pk)
|
||||
else:
|
||||
entry[f"pass@{k}"] = None
|
||||
output[task_id] = entry
|
||||
|
||||
summary = {}
|
||||
for k in k_values:
|
||||
vals = scores[k]
|
||||
if vals:
|
||||
summary[f"pass@{k}"] = round(float(np.mean(vals)), 4)
|
||||
else:
|
||||
summary[f"pass@{k}"] = None
|
||||
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(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,
|
||||
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,
|
||||
help="Specific problem indices (0-based)",
|
||||
help="Skip generation, test existing completions JSON",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
_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(
|
||||
"--generate_only", action="store_true", help="Only generate, skip testing"
|
||||
)
|
||||
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(
|
||||
engine=engine,
|
||||
problems=problems,
|
||||
return EvalConfig(
|
||||
param_path=args.param_path,
|
||||
data_path=args.data_path,
|
||||
output=args.output,
|
||||
test_only=args.test_only,
|
||||
generate_only=args.generate_only,
|
||||
num_samples=args.num_samples,
|
||||
max_tokens=args.max_tokens,
|
||||
temperature=args.temperature,
|
||||
top_p=args.top_p,
|
||||
top_k=args.top_k,
|
||||
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}")
|
||||
for k, v in summary.items():
|
||||
print(f" {k}: {v:.2%}")
|
||||
if v is not None:
|
||||
print(f" {k}: {v:.2%}")
|
||||
else:
|
||||
print(f" {k}: N/A")
|
||||
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__":
|
||||
|
||||
+433
-220
@@ -1,248 +1,397 @@
|
||||
"""IFD (Instruction Following Difficulty) data quality scoring.
|
||||
|
||||
Computes IFD scores for instruction-response pairs to guide data selection.
|
||||
IFD = conditional_NLL / unconditional_NLL, where:
|
||||
IFD = conditional_NLL / unconditional_NLL
|
||||
|
||||
- conditional_NLL: average CE loss on response tokens given instruction context
|
||||
- unconditional_NLL: average CE loss on response tokens alone
|
||||
- Messages format: plain text concatenation (no chat template)
|
||||
- Plain format: raw instr_key + resp_key fields
|
||||
|
||||
Higher IFD (close to 1) = instruction provides less help = harder sample.
|
||||
Lower IFD (close to 0) = instruction provides strong guidance = easy sample.
|
||||
IFD > 1 = instruction misleads the model = likely low-quality data.
|
||||
|
||||
Usage::
|
||||
|
||||
python scripts/eval/ifd.py --param_path ./params \
|
||||
--input data.jsonl --output data_with_ifd.jsonl \
|
||||
--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
|
||||
v2 changelog:
|
||||
- Same token set: unconditional pass prefixes resp with a plain-text sentinel
|
||||
(default ``\\n``; use ``--sentinel_text ""`` for bos/pad fallback).
|
||||
Both branches predict the identical N resp tokens.
|
||||
Single-token answers (rl=1) are now supported.
|
||||
- ctx_len tracked in output
|
||||
- skip_reason for None samples (no more silent None)
|
||||
- --per_token for per-token IFD breakdown
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import statistics
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import tqdm
|
||||
|
||||
from astrai.model import AutoModel
|
||||
from astrai.preprocessing.packing import plan_bfd
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
|
||||
def compute_ifd(
|
||||
model,
|
||||
tokenizer,
|
||||
instruction: str,
|
||||
response: str,
|
||||
device: str,
|
||||
max_len: int = 2048,
|
||||
use_chat_template: bool = False,
|
||||
) -> dict:
|
||||
if use_chat_template:
|
||||
return _compute_ifd_with_template(
|
||||
model, tokenizer, instruction, response, device, max_len
|
||||
def _pack_bins(pairs, max_len):
|
||||
"""BFD bin packing: pack (c+r) into bins of max total length.
|
||||
|
||||
Reuses :func:`plan_bfd` so the BFD heuristic stays single-sourced.
|
||||
"""
|
||||
# Treat each pair as a single sequence of length len(c)+len(r) for
|
||||
# planning purposes; plan_bfd works on pure lengths.
|
||||
fake_sequences = [[0] * (len(c) + len(r)) for c, r in pairs]
|
||||
plan = plan_bfd(fake_sequences, max_len)
|
||||
return [
|
||||
[(i, pairs[i][0], pairs[i][1]) for i in bin_indices] for bin_indices in plan
|
||||
]
|
||||
|
||||
|
||||
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]
|
||||
|
||||
|
||||
def _collect_input_files(input_path: str) -> list:
|
||||
"""Resolve *input_path* to a list of JSONL/JSON files."""
|
||||
if os.path.isdir(input_path):
|
||||
files = []
|
||||
for ext in ("*.jsonl", "*.json"):
|
||||
files.extend(
|
||||
sorted(glob.glob(os.path.join(input_path, "**", ext), recursive=True))
|
||||
)
|
||||
return files
|
||||
return sorted(glob.glob(input_path))
|
||||
|
||||
|
||||
def _load_items(filepath: str) -> list:
|
||||
"""Load JSONL or JSON (array / single dict) into a list of dicts."""
|
||||
with open(filepath, "r", encoding="utf-8") as f:
|
||||
if filepath.lower().endswith(".json"):
|
||||
data = json.load(f)
|
||||
if isinstance(data, dict):
|
||||
return [data]
|
||||
return data
|
||||
return [json.loads(line) for line in f if line.strip()]
|
||||
|
||||
|
||||
@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)
|
||||
)
|
||||
return _compute_ifd_raw(model, tokenizer, instruction, response, device, max_len)
|
||||
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,
|
||||
}
|
||||
|
||||
def _compute_ifd_raw(model, tokenizer, instruction, response, device, max_len) -> dict:
|
||||
instr_ids = tokenizer.encode(instruction)
|
||||
resp_ids = tokenizer.encode(response)
|
||||
# ---- 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
|
||||
|
||||
if not resp_ids:
|
||||
return {
|
||||
"L_cond": None,
|
||||
"L_uncond": None,
|
||||
"ifd": None,
|
||||
"error": "empty response",
|
||||
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
|
||||
|
||||
qa_len = len(instr_ids) + len(resp_ids)
|
||||
if qa_len > max_len:
|
||||
overflow = qa_len - max_len
|
||||
instr_ids = instr_ids[overflow:]
|
||||
|
||||
instr_len = len(instr_ids)
|
||||
resp_len = len(resp_ids)
|
||||
|
||||
qa_ids = instr_ids + resp_ids
|
||||
qa_tensor = torch.tensor([qa_ids], device=device, dtype=torch.long)
|
||||
|
||||
with torch.inference_mode():
|
||||
logits_qa = model(qa_tensor)["logits"][0]
|
||||
|
||||
resp_logits = logits_qa[instr_len - 1 : -1]
|
||||
resp_targets = torch.tensor(resp_ids, 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 {
|
||||
"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,
|
||||
}
|
||||
return result
|
||||
|
||||
|
||||
def _compute_ifd_with_template(
|
||||
model, tokenizer, instruction, response, device, max_len
|
||||
) -> dict:
|
||||
instr_prefix = tokenizer.apply_chat_template(
|
||||
[{"role": "user", "content": instruction}],
|
||||
tokenize=False,
|
||||
add_generation_prompt=True,
|
||||
)
|
||||
full_text = tokenizer.apply_chat_template(
|
||||
[
|
||||
{"role": "user", "content": instruction},
|
||||
{"role": "assistant", "content": response},
|
||||
],
|
||||
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 {
|
||||
"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": prefix_len,
|
||||
"resp_len": len(resp_ids),
|
||||
"error": None,
|
||||
}
|
||||
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 process_file(
|
||||
param_path: str,
|
||||
input_file: str,
|
||||
output_file: str,
|
||||
instr_key: str,
|
||||
resp_key: str,
|
||||
max_len: int,
|
||||
use_chat_template: bool = False,
|
||||
model,
|
||||
tokenizer,
|
||||
input_file,
|
||||
output_file,
|
||||
instr_key,
|
||||
resp_key,
|
||||
max_len=2048,
|
||||
data_format="plain",
|
||||
batch_size=1,
|
||||
device=None,
|
||||
sentinel_ids=None,
|
||||
per_token=False,
|
||||
max_samples=None,
|
||||
):
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
dtype = torch.bfloat16 if device == "cuda" else torch.float32
|
||||
"""Score a single file, write per-sample JSONL, return summary stats."""
|
||||
if device is None:
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
model = AutoModel.from_pretrained(param_path)
|
||||
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
||||
model.to(device=device, dtype=dtype)
|
||||
model.eval()
|
||||
if sentinel_ids is None:
|
||||
sentinel_ids = _resolve_sentinel_ids(tokenizer, "\n")
|
||||
|
||||
if use_chat_template and tokenizer._chat_template is None:
|
||||
raise RuntimeError(
|
||||
"--use_chat_template specified but tokenizer has no chat template. "
|
||||
"Add a chat_template to tokenizer_config.json or omit the flag."
|
||||
)
|
||||
data = _load_items(input_file)
|
||||
|
||||
with open(input_file, "r", encoding="utf-8") as f:
|
||||
data = [json.loads(line) for line in f if line.strip()]
|
||||
if max_samples and len(data) > max_samples:
|
||||
import random
|
||||
|
||||
data = random.sample(data, max_samples)
|
||||
|
||||
results = []
|
||||
ifd_values = []
|
||||
all_ifds = []
|
||||
buffer = []
|
||||
|
||||
with torch.inference_mode():
|
||||
for item in tqdm.tqdm(data, desc="Computing IFD", unit="sample"):
|
||||
instruction = item[instr_key]
|
||||
response = item[resp_key]
|
||||
scores = compute_ifd(
|
||||
label = os.path.splitext(os.path.basename(input_file))[0]
|
||||
|
||||
for item in tqdm.tqdm(data, desc=f" {label}", unit="sample", leave=False):
|
||||
if data_format == "messages":
|
||||
turns = []
|
||||
for i, msg in enumerate(item.get("messages", [])):
|
||||
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,
|
||||
tokenizer,
|
||||
instruction,
|
||||
response,
|
||||
device,
|
||||
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:
|
||||
for item in results:
|
||||
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]
|
||||
stats = {
|
||||
"samples": len(data),
|
||||
"valid_ifd": len(valid_ifd),
|
||||
"skipped": len(data) - len(valid_ifd),
|
||||
}
|
||||
if valid_ifd:
|
||||
import statistics
|
||||
stats["mean_ifd"] = statistics.mean(valid_ifd)
|
||||
stats["median_ifd"] = statistics.median(valid_ifd)
|
||||
if len(valid_ifd) > 1:
|
||||
stats["stdev_ifd"] = statistics.stdev(valid_ifd)
|
||||
stats["min_ifd"] = min(valid_ifd)
|
||||
stats["max_ifd"] = max(valid_ifd)
|
||||
|
||||
print(f"\n{'=' * 50}")
|
||||
print(f" Samples: {len(data)}")
|
||||
print(f" Valid IFD: {len(valid_ifd)}")
|
||||
print(f" Mean IFD: {statistics.mean(valid_ifd):.4f}")
|
||||
print(f" Median IFD: {statistics.median(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" [{label}]")
|
||||
print(f"{'=' * 50}")
|
||||
print(f" Samples: {len(data)}")
|
||||
print(f" Valid IFD: {len(valid_ifd)}")
|
||||
print(f" Skipped: {len(data) - len(valid_ifd)}")
|
||||
print(f" Mean IFD: {statistics.mean(valid_ifd):.4f}")
|
||||
print(f" Median IFD: {statistics.median(valid_ifd):.4f}")
|
||||
if len(valid_ifd) > 1:
|
||||
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" Results saved to {output_file}")
|
||||
return stats
|
||||
|
||||
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():
|
||||
@@ -250,43 +399,107 @@ def main():
|
||||
description="Compute IFD scores for instruction-response data"
|
||||
)
|
||||
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("--output", type=str, required=True, help="Output JSONL file")
|
||||
parser.add_argument(
|
||||
"--instr_key",
|
||||
"--input_path",
|
||||
type=str,
|
||||
default="instruction",
|
||||
help="Key for instruction field",
|
||||
required=True,
|
||||
help="Input file, glob pattern, or directory.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--resp_key",
|
||||
"--output_dir",
|
||||
type=str,
|
||||
default="response",
|
||||
help="Key for response field",
|
||||
required=True,
|
||||
help="Directory for output files (summary.json + per-file JSONL).",
|
||||
)
|
||||
parser.add_argument("--max_len", type=int, default=2048, help="Max token length")
|
||||
parser.add_argument(
|
||||
"--format",
|
||||
type=str,
|
||||
default="plain",
|
||||
choices=["plain", "messages"],
|
||||
help="Input format",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_len",
|
||||
type=int,
|
||||
default=2048,
|
||||
help="Max token length (instruction truncated to fit)",
|
||||
"--instr_key", type=str, default="instruction", help="Key for instruction field"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--no_chat_template",
|
||||
"--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(
|
||||
"--dtype",
|
||||
type=str,
|
||||
default="bfloat16" if torch.cuda.is_available() else "float32",
|
||||
help="Torch dtype",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--sentinel_text",
|
||||
type=str,
|
||||
default="\n",
|
||||
help='Plain-text prefix for unconditional pass (default: "\\n"). Use "" for bos/pad fallback.',
|
||||
)
|
||||
parser.add_argument(
|
||||
"--per_token",
|
||||
action="store_true",
|
||||
default=False,
|
||||
help="Disable chat template, use raw text concatenation",
|
||||
help="Include per-token IFD breakdown in output",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_samples",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Maximum number of samples per file (random subsample). Default: all.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
process_file(
|
||||
args.param_path,
|
||||
args.input,
|
||||
args.output,
|
||||
args.instr_key,
|
||||
args.resp_key,
|
||||
args.max_len,
|
||||
use_chat_template=not args.no_chat_template,
|
||||
)
|
||||
if args.device is None:
|
||||
args.device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
dtype = getattr(torch, args.dtype)
|
||||
|
||||
print(f"Loading model from {args.param_path} ...")
|
||||
model = AutoModel.from_pretrained(args.param_path)
|
||||
tokenizer = AutoTokenizer.from_pretrained(args.param_path)
|
||||
model.to(device=args.device, dtype=dtype)
|
||||
model.eval()
|
||||
|
||||
sentinel_ids = _resolve_sentinel_ids(tokenizer, args.sentinel_text)
|
||||
|
||||
input_files = _collect_input_files(args.input_path)
|
||||
if not input_files:
|
||||
print(f"No input files found at {args.input_path}")
|
||||
return
|
||||
|
||||
print(f"Found {len(input_files)} file(s) to evaluate")
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
all_stats = {}
|
||||
for filepath in input_files:
|
||||
label = os.path.splitext(os.path.basename(filepath))[0]
|
||||
output_file = os.path.join(args.output_dir, f"{label}_ifd.jsonl")
|
||||
|
||||
stats = process_file(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
input_file=filepath,
|
||||
output_file=output_file,
|
||||
instr_key=args.instr_key,
|
||||
resp_key=args.resp_key,
|
||||
max_len=args.max_len,
|
||||
data_format=args.format,
|
||||
batch_size=args.batch_size,
|
||||
device=args.device,
|
||||
sentinel_ids=sentinel_ids,
|
||||
per_token=args.per_token,
|
||||
max_samples=args.max_samples,
|
||||
)
|
||||
all_stats[label] = stats
|
||||
|
||||
summary_path = os.path.join(args.output_dir, "summary.json")
|
||||
with open(summary_path, "w", encoding="utf-8") as f:
|
||||
json.dump(all_stats, f, ensure_ascii=False, indent=2)
|
||||
print(f"\nSummary saved to {summary_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -5,7 +5,7 @@ Supports all IFEval constraint types except language detection.
|
||||
|
||||
Usage::
|
||||
|
||||
python scripts/tools/evaluate_ifeval.py --param_path ./params \
|
||||
python scripts/eval/evaluate_ifeval.py --param_path ./params \
|
||||
--data_path ifeval.jsonl --output results.json \
|
||||
--temperature 0.1 --max_tokens 512
|
||||
"""
|
||||
@@ -14,21 +14,17 @@ import argparse
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import urllib.request
|
||||
from typing import Callable, Dict, List, Optional
|
||||
|
||||
import torch
|
||||
import tqdm
|
||||
from datasets import load_dataset
|
||||
|
||||
from astrai.inference import InferenceEngine
|
||||
from astrai.model import AutoModel
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
IFEVAL_URL = (
|
||||
"https://raw.githubusercontent.com/google-research/"
|
||||
"google-research/master/instruction_following_eval/data/input_data.jsonl"
|
||||
)
|
||||
|
||||
IFEVAL_HF_DATASET = "google/IFEval"
|
||||
CONSTRAINT_VERIFIERS: Dict[str, Callable[[str, dict], bool]] = {}
|
||||
|
||||
|
||||
@@ -310,15 +306,12 @@ def download_ifeval(data_path: str):
|
||||
if os.path.exists(data_path):
|
||||
return
|
||||
os.makedirs(os.path.dirname(data_path) or ".", exist_ok=True)
|
||||
print(f"Downloading IFEval from {IFEVAL_URL} ...")
|
||||
tmp = data_path + ".tmp"
|
||||
urllib.request.urlretrieve(IFEVAL_URL, tmp)
|
||||
with open(tmp, "rb") as f_in:
|
||||
content = f_in.read()
|
||||
with open(data_path, "wb") as f_out:
|
||||
f_out.write(content)
|
||||
os.remove(tmp)
|
||||
print(f" saved to {data_path}")
|
||||
print(f"Downloading IFEval from HuggingFace ({IFEVAL_HF_DATASET}) ...")
|
||||
ds = load_dataset(IFEVAL_HF_DATASET, split="train")
|
||||
with open(data_path, "w", encoding="utf-8") as f:
|
||||
for item in ds:
|
||||
f.write(json.dumps(item, ensure_ascii=False) + "\n")
|
||||
print(f" saved {len(ds)} items to {data_path}")
|
||||
|
||||
|
||||
def load_problems(data_path: str) -> List[dict]:
|
||||
@@ -571,7 +564,7 @@ def main():
|
||||
print(f" Unsupported: {summary['unsupported_constraints']}")
|
||||
print(f"{'=' * 60}")
|
||||
|
||||
print(f"\nPer-type accuracy:")
|
||||
print("\nPer-type accuracy:")
|
||||
for inst_id, stats in sorted(summary["per_type_accuracy"].items()):
|
||||
print(
|
||||
f" {inst_id:50s} {stats['accuracy']:.2%} "
|
||||
|
||||
@@ -4,18 +4,18 @@ import argparse
|
||||
import csv
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
import tarfile
|
||||
import random
|
||||
from collections import defaultdict
|
||||
|
||||
import requests
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import tqdm
|
||||
from datasets import load_dataset
|
||||
|
||||
from astrai.model import AutoModel
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
MMLU_URL = "https://people.eecs.berkeley.edu/~hendrycks/data.tar"
|
||||
MMLU_HF_DATASET = "cais/mmlu"
|
||||
MMLU_SUBJECTS = [
|
||||
"abstract_algebra",
|
||||
"anatomy",
|
||||
@@ -77,38 +77,40 @@ MMLU_SUBJECTS = [
|
||||
]
|
||||
|
||||
|
||||
def _download_and_extract(url: str, data_dir: str):
|
||||
tar_path = os.path.join(data_dir, "data.tar")
|
||||
os.makedirs(data_dir, exist_ok=True)
|
||||
print(f"Downloading MMLU data from {url}...")
|
||||
resp = requests.get(url, stream=True, timeout=300)
|
||||
resp.raise_for_status()
|
||||
total = int(resp.headers.get("content-length", 0))
|
||||
with tqdm.tqdm(total=total, unit="B", unit_scale=True, desc=" Download") as bar:
|
||||
with open(tar_path, "wb") as f:
|
||||
for chunk in resp.iter_content(chunk_size=8192):
|
||||
f.write(chunk)
|
||||
bar.update(len(chunk))
|
||||
print("Extracting...")
|
||||
with tarfile.open(tar_path, "r") as tf:
|
||||
tf.extractall(data_dir)
|
||||
os.remove(tar_path)
|
||||
def _write_subject_csv(data_dir: str, split: str, subject: str, rows: list[dict]):
|
||||
split_dir = os.path.join(data_dir, split)
|
||||
os.makedirs(split_dir, exist_ok=True)
|
||||
path = os.path.join(split_dir, f"{subject}_{split}.csv")
|
||||
with open(path, "w", encoding="utf-8", newline="") as f:
|
||||
writer = csv.writer(f)
|
||||
for row in rows:
|
||||
writer.writerow(row)
|
||||
|
||||
|
||||
def download_mmlu(data_dir: str):
|
||||
_download_and_extract(MMLU_URL, data_dir)
|
||||
src = os.path.join(data_dir, "data")
|
||||
if os.path.exists(src):
|
||||
for item in os.listdir(src):
|
||||
src_item = os.path.join(src, item)
|
||||
dst_item = os.path.join(data_dir, item)
|
||||
if os.path.exists(dst_item):
|
||||
if os.path.isdir(dst_item):
|
||||
shutil.rmtree(dst_item)
|
||||
else:
|
||||
os.remove(dst_item)
|
||||
os.rename(src_item, dst_item)
|
||||
os.rmdir(src)
|
||||
print(f"Downloading MMLU from HuggingFace ({MMLU_HF_DATASET}) ...")
|
||||
letters = ("A", "B", "C", "D")
|
||||
split_map = {"dev": "dev", "val": "validation", "test": "test"}
|
||||
for local_split, hf_split in split_map.items():
|
||||
ds = load_dataset(MMLU_HF_DATASET, "all", split=hf_split)
|
||||
grouped: dict[str, list[dict]] = defaultdict(list)
|
||||
for item in tqdm.tqdm(ds, desc=f" {local_split}", leave=False):
|
||||
subject = item["subject"]
|
||||
choices = item["choices"]
|
||||
ans_letter = letters[item["answer"]]
|
||||
grouped[subject].append(
|
||||
[
|
||||
item["question"],
|
||||
f"A){choices[0]}",
|
||||
f"B){choices[1]}",
|
||||
f"C){choices[2]}",
|
||||
f"D){choices[3]}",
|
||||
ans_letter,
|
||||
]
|
||||
)
|
||||
for subject, rows in grouped.items():
|
||||
_write_subject_csv(data_dir, local_split, subject, rows)
|
||||
print(f" {local_split}: {len(ds)} items, {len(grouped)} subjects")
|
||||
print(f"MMLU data saved to {data_dir}")
|
||||
|
||||
|
||||
@@ -139,17 +141,12 @@ def load_csv(path: str) -> list[dict]:
|
||||
return data
|
||||
|
||||
|
||||
def build_prompt(
|
||||
question: str, choices: dict, subject: str, n_shot: int, dev_data: list[dict]
|
||||
) -> str:
|
||||
prompt = ""
|
||||
if n_shot > 0 and dev_data:
|
||||
prompt = f"The following are multiple choice questions (with answers) about {subject}.\n\n"
|
||||
for item in dev_data[:n_shot]:
|
||||
prompt += f"Question: {item['question']}\n"
|
||||
for k in ("A", "B", "C", "D"):
|
||||
prompt += f"{k}. {item[k]}\n"
|
||||
prompt += f"Answer: {item['answer']}\n\n"
|
||||
def build_prompt(question: str, choices: dict, subject: str) -> str:
|
||||
"""Build the raw question prompt (without few-shot examples).
|
||||
|
||||
Few-shot examples are handled by ``apply_chat`` to avoid duplication.
|
||||
"""
|
||||
prompt = f"The following are multiple choice questions (with answers) about {subject}.\n\n"
|
||||
prompt += f"Question: {question}\n"
|
||||
for k in ("A", "B", "C", "D"):
|
||||
prompt += f"{k}. {choices[k]}\n"
|
||||
@@ -158,19 +155,22 @@ def build_prompt(
|
||||
|
||||
|
||||
def apply_chat(
|
||||
tokenizer, raw_prompt: str, n_shot: int, dev_data: list[dict] | None
|
||||
tokenizer,
|
||||
raw_prompt: str,
|
||||
n_shot: int,
|
||||
dev_data: list[dict] | None,
|
||||
subject: str = "",
|
||||
) -> str:
|
||||
"""Wrap raw MMLU prompt in the model's chat template format.
|
||||
|
||||
For few-shot, prepend example Q&A pairs as a second user/assistant exchange.
|
||||
For few-shot, prepend example Q&A pairs as user/assistant exchanges.
|
||||
Few-shot examples use the same subject preamble as the test question to
|
||||
keep the format consistent.
|
||||
"""
|
||||
messages = []
|
||||
if n_shot > 0 and dev_data:
|
||||
for item in dev_data[:n_shot]:
|
||||
q = f"Question: {item['question']}\n"
|
||||
for k in ("A", "B", "C", "D"):
|
||||
q += f"{k}. {item[k]}\n"
|
||||
q += "Answer:"
|
||||
q = build_prompt(item["question"], item, subject)
|
||||
messages.append({"role": "user", "content": q})
|
||||
messages.append({"role": "assistant", "content": item["answer"]})
|
||||
messages.append({"role": "user", "content": raw_prompt})
|
||||
@@ -206,6 +206,25 @@ def choice_logprob(
|
||||
return score
|
||||
|
||||
|
||||
def _permute_choices(item: dict, rng: random.Random) -> tuple[dict, str]:
|
||||
"""Shuffle the option order of a question.
|
||||
|
||||
Returns ``(permuted_item, new_answer_letter)``. The question text and
|
||||
the *content* of each choice are unchanged; only which letter (A/B/C/D)
|
||||
maps to which content is shuffled. This neutralises the model's
|
||||
positional bias (e.g. always picking B).
|
||||
"""
|
||||
letters = ("A", "B", "C", "D")
|
||||
contents = [item[k] for k in letters]
|
||||
perm = list(letters)
|
||||
rng.shuffle(perm)
|
||||
permuted = {"question": item["question"]}
|
||||
for new_letter, orig_letter in zip(letters, perm):
|
||||
permuted[new_letter] = item[orig_letter]
|
||||
new_answer = letters[perm.index(item["answer"])]
|
||||
return permuted, new_answer
|
||||
|
||||
|
||||
def evaluate_subject(
|
||||
model,
|
||||
tokenizer,
|
||||
@@ -214,20 +233,24 @@ def evaluate_subject(
|
||||
dev_data: list[dict] | None,
|
||||
device: str,
|
||||
n_shot: int,
|
||||
seed: int = 0,
|
||||
) -> tuple[float, int, int]:
|
||||
rng = random.Random(seed) if seed >= 0 else None
|
||||
correct = 0
|
||||
total = 0
|
||||
for item in tqdm.tqdm(test_data, desc=f"{subject:40s}", leave=False):
|
||||
raw_prompt = build_prompt(
|
||||
item["question"], item, subject, n_shot, dev_data or []
|
||||
)
|
||||
context = apply_chat(tokenizer, raw_prompt, n_shot, dev_data or [])
|
||||
if rng is not None:
|
||||
permuted, answer = _permute_choices(item, rng)
|
||||
else:
|
||||
permuted, answer = item, item["answer"]
|
||||
raw_prompt = build_prompt(permuted["question"], permuted, subject)
|
||||
context = apply_chat(tokenizer, raw_prompt, n_shot, dev_data or [], subject)
|
||||
context_ids = tokenizer.encode(context)
|
||||
scores = {
|
||||
c: choice_logprob(model, tokenizer, context_ids, c, device)
|
||||
for c in ("A", "B", "C", "D")
|
||||
}
|
||||
if max(scores, key=scores.get) == item["answer"]:
|
||||
if max(scores, key=scores.get) == answer:
|
||||
correct += 1
|
||||
total += 1
|
||||
return correct / total, correct, total
|
||||
@@ -262,6 +285,12 @@ def main():
|
||||
default="bfloat16" if torch.cuda.is_available() else "float32",
|
||||
help="Torch dtype",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--seed",
|
||||
type=int,
|
||||
default=0,
|
||||
help="Seed for option permutation (0 to enable, -1 to disable)",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.download or not os.path.exists(args.data_dir):
|
||||
@@ -293,7 +322,14 @@ def main():
|
||||
test_data = load_csv(test_path)
|
||||
|
||||
acc, corr, tot = evaluate_subject(
|
||||
model, tokenizer, subject, test_data, dev_data, device, args.n_shot
|
||||
model,
|
||||
tokenizer,
|
||||
subject,
|
||||
test_data,
|
||||
dev_data,
|
||||
device,
|
||||
args.n_shot,
|
||||
seed=args.seed,
|
||||
)
|
||||
results[subject] = {"accuracy": round(acc, 4), "correct": corr, "total": tot}
|
||||
total_correct += corr
|
||||
|
||||
+422
-69
@@ -1,5 +1,9 @@
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import statistics
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
@@ -9,95 +13,400 @@ from astrai.model import AutoModel
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
|
||||
def _collect_input_files(input_path: str) -> List[str]:
|
||||
"""Resolve *input_path* to a list of JSONL/JSON files."""
|
||||
if os.path.isdir(input_path):
|
||||
files = []
|
||||
for ext in ("*.jsonl", "*.json"):
|
||||
files.extend(
|
||||
sorted(glob.glob(os.path.join(input_path, "**", ext), recursive=True))
|
||||
)
|
||||
return files
|
||||
return sorted(glob.glob(input_path))
|
||||
|
||||
|
||||
def _load_items(filepath: str) -> List[dict]:
|
||||
"""Load JSONL or JSON (array / single dict) into a list of dicts."""
|
||||
with open(filepath, "r", encoding="utf-8") as f:
|
||||
if filepath.lower().endswith(".json"):
|
||||
data = json.load(f)
|
||||
if isinstance(data, dict):
|
||||
return [data]
|
||||
return data
|
||||
return [json.loads(line) for line in f if line.strip()]
|
||||
|
||||
|
||||
def _encode_batch(
|
||||
tokenizer: AutoTokenizer, texts: List[str], max_length: int
|
||||
) -> Tuple[List[List[int]], List[List[int]]]:
|
||||
"""Encode *texts* and return (token_ids, attention_masks).
|
||||
|
||||
Each sequence is left-aligned and padded to the batch max length.
|
||||
"""
|
||||
encoded = [tokenizer.encode(t)[:max_length] for t in texts]
|
||||
if not encoded:
|
||||
return [], []
|
||||
max_len = max(len(seq) for seq in encoded)
|
||||
padded_ids = []
|
||||
masks = []
|
||||
for seq in encoded:
|
||||
pad_len = max_len - len(seq)
|
||||
padded_ids.append(seq + [tokenizer.pad_id] * pad_len)
|
||||
masks.append([1] * len(seq) + [0] * pad_len)
|
||||
return padded_ids, masks
|
||||
|
||||
|
||||
def _compute_batch(
|
||||
model,
|
||||
input_ids: torch.Tensor,
|
||||
attention_mask: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Forward pass and return (log_probs, valid_mask) of shape [B, S-1].
|
||||
|
||||
log_probs[i, j] = log P(token j+1 | tokens 0..j)
|
||||
"""
|
||||
output = model(input_ids, input_mask=attention_mask)
|
||||
logits = output["logits"][:, :-1, :] # [B, S-1, V]
|
||||
targets = input_ids[:, 1:] # [B, S-1]
|
||||
valid = attention_mask[:, 1:].float() # [B, S-1]
|
||||
|
||||
log_probs = F.log_softmax(logits.float(), dim=-1) # [B, S-1, V]
|
||||
token_log_probs = log_probs.gather(2, targets.unsqueeze(-1)).squeeze(-1) # [B, S-1]
|
||||
|
||||
return token_log_probs, valid
|
||||
|
||||
|
||||
def _token_type(token_id: int, stop_ids: frozenset, decode_fn) -> str:
|
||||
"""Classify a token into a coarse type for analysis.
|
||||
|
||||
*stop_ids* is a pre-built set of special token IDs.
|
||||
*decode_fn* is ``tokenizer.decode`` (or a wrapper) for single-token
|
||||
decoding.
|
||||
"""
|
||||
if token_id in stop_ids:
|
||||
return "special"
|
||||
decoded = decode_fn([token_id], skip_special_tokens=True)
|
||||
if any("\u4e00" <= ch <= "\u9fff" for ch in decoded):
|
||||
return "cjk"
|
||||
if any(ord(ch) > 127 for ch in decoded):
|
||||
return "non_ascii"
|
||||
return "ascii"
|
||||
|
||||
|
||||
def _percentiles(values: List[float]) -> Dict[str, float]:
|
||||
"""Compute common percentiles from a list of floats.
|
||||
|
||||
Uses linear interpolation between closest ranks (same convention
|
||||
as NumPy's default).
|
||||
"""
|
||||
if not values:
|
||||
return {}
|
||||
sorted_vals = sorted(values)
|
||||
n = len(sorted_vals)
|
||||
|
||||
def _pct(p: float) -> float:
|
||||
if n == 1:
|
||||
return sorted_vals[0]
|
||||
k = p * (n - 1)
|
||||
f = int(k)
|
||||
c = min(f + 1, n - 1)
|
||||
return sorted_vals[f] + (sorted_vals[c] - sorted_vals[f]) * (k - f)
|
||||
|
||||
return {
|
||||
"p50": _pct(0.50),
|
||||
"p90": _pct(0.90),
|
||||
"p95": _pct(0.95),
|
||||
"p99": _pct(0.99),
|
||||
}
|
||||
|
||||
|
||||
class LossAccumulator:
|
||||
"""Accumulate per-token losses with optional streaming mode.
|
||||
|
||||
When *stream* is True (token_level=False), losses are not kept
|
||||
in memory individually — only a running sum/count and a histogram
|
||||
(for approximate percentiles) are maintained. When *stream* is
|
||||
False, all losses are retained for exact statistics and per-record
|
||||
output.
|
||||
"""
|
||||
|
||||
_HIST_BINS = 1000
|
||||
_HIST_MAX = 20.0 # clamp losses above this for histogram
|
||||
|
||||
def __init__(self, stream: bool):
|
||||
self.stream = stream
|
||||
self.losses: List[float] = [] if not stream else []
|
||||
self.total: float = 0.0
|
||||
self.count: int = 0
|
||||
self.hist = torch.zeros(self._HIST_BINS, dtype=torch.long)
|
||||
# per-type losses (only populated when not streaming)
|
||||
self.by_type: Dict[str, List[float]] = {}
|
||||
self.type_total: Dict[str, float] = {}
|
||||
self.type_count: Dict[str, int] = {}
|
||||
|
||||
def add(self, losses: List[float]):
|
||||
self.total += sum(losses)
|
||||
self.count += len(losses)
|
||||
if self.stream:
|
||||
clamped = [min(max(l, 0.0), self._HIST_MAX) for l in losses]
|
||||
idx = torch.tensor(clamped) / self._HIST_MAX * (self._HIST_BINS - 1)
|
||||
self.hist += torch.bincount(
|
||||
idx.long().clamp(0, self._HIST_BINS - 1),
|
||||
minlength=self._HIST_BINS,
|
||||
)
|
||||
else:
|
||||
self.losses.extend(losses)
|
||||
|
||||
def add_typed(self, ttype: str, losses: List[float]):
|
||||
if not self.stream:
|
||||
self.by_type.setdefault(ttype, []).extend(losses)
|
||||
self.type_total[ttype] = self.type_total.get(ttype, 0.0) + sum(losses)
|
||||
self.type_count[ttype] = self.type_count.get(ttype, 0) + len(losses)
|
||||
|
||||
def stats(self) -> Dict:
|
||||
result: Dict = {}
|
||||
if self.count == 0:
|
||||
return result
|
||||
mean_loss = self.total / self.count
|
||||
result["overall"] = {
|
||||
"num_tokens": self.count,
|
||||
"mean_loss": mean_loss,
|
||||
"ppl": float(torch.exp(torch.tensor(mean_loss))),
|
||||
}
|
||||
if self.stream:
|
||||
result["overall"].update(self._hist_percentiles())
|
||||
else:
|
||||
result["overall"]["median_loss"] = statistics.median(self.losses)
|
||||
result["overall"].update(_percentiles(self.losses))
|
||||
|
||||
if self.type_count:
|
||||
result["by_token_type"] = {}
|
||||
for ttype in sorted(self.type_count.keys()):
|
||||
cnt = self.type_count[ttype]
|
||||
tmean = self.type_total[ttype] / cnt
|
||||
entry: Dict = {
|
||||
"num_tokens": cnt,
|
||||
"mean_loss": tmean,
|
||||
"ppl": float(torch.exp(torch.tensor(tmean))),
|
||||
}
|
||||
if not self.stream and ttype in self.by_type:
|
||||
entry["median_loss"] = statistics.median(self.by_type[ttype])
|
||||
entry.update(_percentiles(self.by_type[ttype]))
|
||||
result["by_token_type"][ttype] = entry
|
||||
return result
|
||||
|
||||
def _hist_percentiles(self) -> Dict[str, float]:
|
||||
"""Approximate percentiles from the histogram."""
|
||||
total = self.hist.sum().item()
|
||||
if total == 0:
|
||||
return {}
|
||||
cum = torch.cumsum(self.hist.float(), dim=0)
|
||||
result = {}
|
||||
for label, p in [("p50", 0.5), ("p90", 0.9), ("p95", 0.95), ("p99", 0.99)]:
|
||||
target = p * total
|
||||
idx = int(torch.searchsorted(cum, target).item())
|
||||
idx = min(idx, self._HIST_BINS - 1)
|
||||
result[label] = (idx + 0.5) / self._HIST_BINS * self._HIST_MAX
|
||||
return result
|
||||
|
||||
|
||||
def process_file(
|
||||
param_path: str, input_file: str, output_file: str, batch_size: int, text_key: str
|
||||
model,
|
||||
tokenizer: AutoTokenizer,
|
||||
items: List[dict],
|
||||
text_key: str,
|
||||
batch_size: int,
|
||||
max_length: int,
|
||||
token_level: bool,
|
||||
max_samples: Optional[int],
|
||||
output_file: Optional[str],
|
||||
label: str,
|
||||
device: str = "cuda",
|
||||
) -> Dict:
|
||||
"""Evaluate a single dataset (list of items), return summary stats.
|
||||
|
||||
If *token_level* is True and *output_file* is set, per-record token_ids
|
||||
and log_probs are written as JSONL alongside the summary.
|
||||
"""
|
||||
if max_samples and len(items) > max_samples:
|
||||
import random
|
||||
|
||||
items = random.sample(items, max_samples)
|
||||
|
||||
texts = [item[text_key] for item in items if text_key in item]
|
||||
print(f" [{label}] {len(texts)} samples, text_key='{text_key}'")
|
||||
|
||||
acc = LossAccumulator(stream=not token_level)
|
||||
per_sample: List[dict] = []
|
||||
|
||||
if token_level:
|
||||
stop_ids = frozenset(tokenizer.stop_ids)
|
||||
decode_fn = tokenizer.decode
|
||||
|
||||
num_batches = (len(texts) + batch_size - 1) // batch_size
|
||||
for i in tqdm.tqdm(
|
||||
range(0, len(texts), batch_size),
|
||||
total=num_batches,
|
||||
desc=f" {label}",
|
||||
leave=False,
|
||||
):
|
||||
batch_texts = texts[i : i + batch_size]
|
||||
padded_ids, masks = _encode_batch(tokenizer, batch_texts, max_length)
|
||||
|
||||
input_ids = torch.tensor(padded_ids, device=device, dtype=torch.long)
|
||||
attention_mask = torch.tensor(masks, device=device, dtype=torch.bool)
|
||||
|
||||
token_log_probs, valid = _compute_batch(model, input_ids, attention_mask)
|
||||
|
||||
for b in range(len(batch_texts)):
|
||||
seq_len = int(valid[b].sum().item())
|
||||
lps = token_log_probs[b, :seq_len].tolist()
|
||||
losses = [-lp for lp in lps]
|
||||
acc.add(losses)
|
||||
|
||||
if token_level:
|
||||
# log_probs correspond to positions 1..seq_len (predicted
|
||||
# from position 0..seq_len-1), so token_ids must skip BOS
|
||||
# at position 0 to stay aligned with log_probs.
|
||||
ids = padded_ids[b][1 : seq_len + 1]
|
||||
per_sample.append(
|
||||
{
|
||||
"text": batch_texts[b][:200],
|
||||
"token_ids": ids,
|
||||
"log_probs": [round(lp, 4) for lp in lps],
|
||||
"ppl": float(torch.exp(torch.tensor(statistics.mean(losses))))
|
||||
if losses
|
||||
else None,
|
||||
}
|
||||
)
|
||||
typed_losses: Dict[str, List[float]] = {}
|
||||
for tid, loss in zip(ids, losses):
|
||||
ttype = _token_type(tid, stop_ids, decode_fn)
|
||||
typed_losses.setdefault(ttype, []).append(loss)
|
||||
for ttype, tl in typed_losses.items():
|
||||
acc.add_typed(ttype, tl)
|
||||
|
||||
stats = acc.stats()
|
||||
|
||||
if token_level and output_file:
|
||||
with open(output_file, "w", encoding="utf-8") as f:
|
||||
for item in per_sample:
|
||||
f.write(json.dumps(item, ensure_ascii=False) + "\n")
|
||||
|
||||
return stats
|
||||
|
||||
|
||||
def print_stats(label: str, stats: Dict):
|
||||
"""Pretty-print summary statistics."""
|
||||
print(f"\n{'=' * 60}")
|
||||
print(f" {label}")
|
||||
print(f"{'=' * 60}")
|
||||
ov = stats.get("overall", {})
|
||||
if ov:
|
||||
print(f" tokens: {ov['num_tokens']:,}")
|
||||
print(f" mean loss: {ov['mean_loss']:.4f}")
|
||||
if "median_loss" in ov:
|
||||
print(f" median loss: {ov['median_loss']:.4f}")
|
||||
print(f" ppl: {ov['ppl']:.2f}")
|
||||
if "p50" in ov:
|
||||
print(
|
||||
f" p50/p90/p95/p99: "
|
||||
f"{ov['p50']:.2f} / {ov['p90']:.2f} / {ov['p95']:.2f} / {ov['p99']:.2f}"
|
||||
)
|
||||
by_type = stats.get("by_token_type", {})
|
||||
if by_type:
|
||||
print(f"\n by token type:")
|
||||
print(f" {'type':<12} {'count':>8} {'mean_loss':>10} {'ppl':>8}")
|
||||
print(f" {'-' * 12} {'-' * 8} {'-' * 10} {'-' * 8}")
|
||||
for ttype, s in by_type.items():
|
||||
print(
|
||||
f" {ttype:<12} {s['num_tokens']:>8,} "
|
||||
f"{s['mean_loss']:>10.4f} {s['ppl']:>8.2f}"
|
||||
)
|
||||
|
||||
|
||||
def main(
|
||||
param_path: str,
|
||||
input_path: str,
|
||||
output_dir: str,
|
||||
text_key: str,
|
||||
batch_size: int,
|
||||
max_length: int,
|
||||
token_level: bool,
|
||||
max_samples: Optional[int],
|
||||
device: str = "cuda",
|
||||
dtype: str = "bfloat16",
|
||||
):
|
||||
# Load model and tokenizer
|
||||
print(f"Loading model from {param_path} ...")
|
||||
model = AutoModel.from_pretrained(param_path)
|
||||
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
||||
model.to(device="cuda", dtype=torch.bfloat16)
|
||||
torch_dtype = getattr(torch, dtype)
|
||||
model.to(device=device, dtype=torch_dtype)
|
||||
model.eval()
|
||||
|
||||
with open(input_file, "r", encoding="utf-8") as f:
|
||||
input_data = [json.loads(line) for line in f]
|
||||
input_files = _collect_input_files(input_path)
|
||||
if not input_files:
|
||||
print(f"No input files found at {input_path}")
|
||||
return
|
||||
|
||||
texts = [item[text_key] for item in input_data]
|
||||
print(f"Found {len(input_files)} file(s) to evaluate")
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
# Encode all texts
|
||||
print(f"Encoding {len(texts)} texts...")
|
||||
encoded_texts = [tokenizer.encode(text) for text in texts]
|
||||
all_stats = {}
|
||||
for filepath in input_files:
|
||||
label = os.path.splitext(os.path.basename(filepath))[0]
|
||||
items = _load_items(filepath)
|
||||
if not items:
|
||||
print(f" [{label}] empty, skipping")
|
||||
continue
|
||||
|
||||
output_data = []
|
||||
total_batches = (len(encoded_texts) + batch_size - 1) // batch_size
|
||||
|
||||
for i in tqdm.tqdm(
|
||||
range(0, len(encoded_texts), batch_size),
|
||||
total=total_batches,
|
||||
desc="Computing perplexity",
|
||||
):
|
||||
batch_encoded = encoded_texts[i : i + batch_size]
|
||||
batch_texts = texts[i : i + batch_size]
|
||||
|
||||
# Find max length in batch and pad
|
||||
max_len = max(len(seq) for seq in batch_encoded)
|
||||
padded_ids = []
|
||||
masks = []
|
||||
|
||||
for seq in batch_encoded:
|
||||
pad_len = max_len - len(seq)
|
||||
padded_seq = seq + [tokenizer.pad_id] * pad_len
|
||||
mask = [True] * len(seq) + [False] * pad_len
|
||||
padded_ids.append(padded_seq)
|
||||
masks.append(mask)
|
||||
|
||||
# Convert to tensors
|
||||
input_ids = torch.tensor(padded_ids, device="cuda", dtype=torch.long)
|
||||
input_mask = torch.tensor(masks, device="cuda", dtype=torch.bool)
|
||||
|
||||
# Compute perplexity
|
||||
output = model(input_ids, input_mask=input_mask)
|
||||
logits = output["logits"]
|
||||
|
||||
# Shift for causal language modeling
|
||||
shifted_logits = logits[:, :-1, :] # [batch_size, seq_len-1, vocab_size]
|
||||
shifted_input_ids = input_ids[:, 1:] # [batch_size, seq_len-1]
|
||||
shifted_mask = input_mask[:, 1:] # [batch_size, seq_len-1]
|
||||
|
||||
# Compute cross entropy loss
|
||||
loss = F.cross_entropy(
|
||||
shifted_logits.flatten(0, 1),
|
||||
shifted_input_ids.flatten(0, 1),
|
||||
reduction="none",
|
||||
token_output = (
|
||||
os.path.join(output_dir, f"{label}_tokens.jsonl") if token_level else None
|
||||
)
|
||||
|
||||
loss = loss.view(shifted_input_ids.shape) # [batch_size, seq_len-1]
|
||||
loss = loss * shifted_mask
|
||||
sentence_loss = loss.sum(dim=1) / shifted_mask.sum(dim=1).clamp(min=1)
|
||||
perplexity = torch.exp(sentence_loss) # [batch_size]
|
||||
stats = process_file(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
items=items,
|
||||
text_key=text_key,
|
||||
batch_size=batch_size,
|
||||
max_length=max_length,
|
||||
token_level=token_level,
|
||||
max_samples=max_samples,
|
||||
output_file=token_output,
|
||||
label=label,
|
||||
device=device,
|
||||
)
|
||||
all_stats[label] = stats
|
||||
print_stats(label, stats)
|
||||
|
||||
for text, ppl in zip(batch_texts, perplexity):
|
||||
output_data.append({text_key: text, "ppl": float(ppl.item())})
|
||||
if token_output:
|
||||
print(f" token-level output: {token_output}")
|
||||
|
||||
# Write results
|
||||
with open(output_file, "w", encoding="utf-8") as f:
|
||||
for item in output_data:
|
||||
f.write(json.dumps(item, ensure_ascii=False) + "\n")
|
||||
|
||||
print(f"Perplexity computation complete. Results saved to {output_file}")
|
||||
summary_path = os.path.join(output_dir, "summary.json")
|
||||
with open(summary_path, "w", encoding="utf-8") as f:
|
||||
json.dump(all_stats, f, ensure_ascii=False, indent=2)
|
||||
print(f"\nSummary saved to {summary_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="Run perplexity with a Khaosz model.")
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Perplexity and token-level loss evaluation on JSONL/JSON data."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--param_path", type=str, required=True, help="Path to the model directory."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--input_file", type=str, required=True, help="Path to the input file."
|
||||
"--input_path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to input file, glob pattern, or directory.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output_file", type=str, required=True, help="Path to the output file."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size", type=int, default=4, help="Batch size for evaluation."
|
||||
"--output_dir",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Directory for output files (summary.json + per-file token JSONL).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text_key",
|
||||
@@ -105,7 +414,51 @@ if __name__ == "__main__":
|
||||
default="text",
|
||||
help="Key for the text field in the input data.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size", type=int, default=4, help="Batch size for evaluation."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_length",
|
||||
type=int,
|
||||
default=2048,
|
||||
help="Maximum sequence length (tokens). Longer sequences are truncated.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--token_level",
|
||||
action="store_true",
|
||||
help="Store per-token log_probs and token type analysis. "
|
||||
"Default: off (only aggregate stats).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_samples",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Maximum number of samples per file (random subsample). Default: all.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--device",
|
||||
type=str,
|
||||
default="cuda" if torch.cuda.is_available() else "cpu",
|
||||
help="Device for model inference.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dtype",
|
||||
type=str,
|
||||
default="bfloat16" if torch.cuda.is_available() else "float32",
|
||||
help="Torch dtype for model weights.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
with torch.inference_mode():
|
||||
process_file(**vars(args))
|
||||
main(
|
||||
param_path=args.param_path,
|
||||
input_path=args.input_path,
|
||||
output_dir=args.output_dir,
|
||||
text_key=args.text_key,
|
||||
batch_size=args.batch_size,
|
||||
max_length=args.max_length,
|
||||
token_level=args.token_level,
|
||||
max_samples=args.max_samples,
|
||||
device=args.device,
|
||||
dtype=args.dtype,
|
||||
)
|
||||
|
||||
@@ -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"""
|
||||
|
||||
import argparse
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.config import AutoRegressiveLMConfig
|
||||
from astrai.inference import KVCache
|
||||
from astrai.inference import ContiguousCache, PageCache
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
|
||||
|
||||
@@ -24,41 +25,14 @@ class GenerationBenchmark:
|
||||
config: AutoRegressiveLMConfig,
|
||||
device: str = "cuda",
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
page_size: int = 128,
|
||||
cache_type: str = "contiguous",
|
||||
):
|
||||
self.config = config
|
||||
self.device = device
|
||||
self.dtype = dtype
|
||||
self.cache_type = cache_type
|
||||
self.model = AutoRegressiveLM(config).to(device=device, dtype=dtype)
|
||||
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()
|
||||
def run_prefill_benchmark(
|
||||
@@ -68,8 +42,12 @@ class GenerationBenchmark:
|
||||
num_trials: int = 10,
|
||||
) -> BenchmarkResult:
|
||||
for _ in range(3):
|
||||
prompt_ids, _ = self._prepare_inputs(
|
||||
batch_size, prompt_length, prompt_length
|
||||
prompt_ids = torch.randint(
|
||||
0,
|
||||
self.config.vocab_size,
|
||||
(batch_size, prompt_length),
|
||||
device=self.device,
|
||||
dtype=torch.long,
|
||||
)
|
||||
_ = self.model(prompt_ids)
|
||||
torch.cuda.synchronize()
|
||||
@@ -78,12 +56,15 @@ class GenerationBenchmark:
|
||||
total_tokens = batch_size * prompt_length * num_trials
|
||||
|
||||
for trial in range(num_trials):
|
||||
prompt_ids, _ = self._prepare_inputs(
|
||||
batch_size, prompt_length, prompt_length
|
||||
prompt_ids = torch.randint(
|
||||
0,
|
||||
self.config.vocab_size,
|
||||
(batch_size, prompt_length),
|
||||
device=self.device,
|
||||
dtype=torch.long,
|
||||
)
|
||||
start = torch.cuda.Event(enable_timing=True)
|
||||
end = torch.cuda.Event(enable_timing=True)
|
||||
|
||||
start.record()
|
||||
_ = self.model(prompt_ids)
|
||||
end.record()
|
||||
@@ -107,6 +88,7 @@ class GenerationBenchmark:
|
||||
"prompt_length": prompt_length,
|
||||
"dtype": str(self.dtype),
|
||||
"device": self.device,
|
||||
"cache": "none",
|
||||
},
|
||||
)
|
||||
|
||||
@@ -120,29 +102,56 @@ class GenerationBenchmark:
|
||||
) -> BenchmarkResult:
|
||||
total_time = 0.0
|
||||
total_tokens = batch_size * gen_length * num_trials
|
||||
page_size = self._page_cache.page_size
|
||||
|
||||
for trial in range(num_trials):
|
||||
prompt_ids, gen_ids = self._prepare_inputs(
|
||||
batch_size,
|
||||
prompt_length,
|
||||
prompt_length + gen_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,
|
||||
prompt_ids = torch.randint(
|
||||
0,
|
||||
self.config.vocab_size,
|
||||
(batch_size, prompt_length),
|
||||
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(
|
||||
prompt_ids,
|
||||
paged_cache=cv,
|
||||
@@ -152,37 +161,35 @@ class GenerationBenchmark:
|
||||
.unsqueeze(0)
|
||||
.expand(batch_size, -1),
|
||||
)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
|
||||
start = torch.cuda.Event(enable_timing=True)
|
||||
end = torch.cuda.Event(enable_timing=True)
|
||||
|
||||
start.record()
|
||||
current_pos = prompt_length
|
||||
|
||||
for i in range(gen_length):
|
||||
input_token = gen_ids[:, i : i + 1]
|
||||
cv = self._page_cache.bind(page_table, total_len=current_pos + 1)
|
||||
pos = prompt_length + i
|
||||
cv = cache.bind_tasks(task_ids, pos + 1, self.device)
|
||||
_ = self.model(
|
||||
input_token,
|
||||
gen_ids[:, i : i + 1],
|
||||
paged_cache=cv,
|
||||
position_ids=torch.full(
|
||||
(batch_size, 1),
|
||||
current_pos,
|
||||
pos,
|
||||
dtype=torch.long,
|
||||
device=self.device,
|
||||
),
|
||||
)
|
||||
current_pos += 1
|
||||
|
||||
end.record()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
for tid in task_ids:
|
||||
cache.task_free(tid)
|
||||
|
||||
trial_time = start.elapsed_time(end) / 1000
|
||||
total_time += trial_time
|
||||
|
||||
for idx in pages:
|
||||
self._page_cache._pool.free(idx)
|
||||
|
||||
print(
|
||||
f" Trial {trial + 1}/{num_trials}: {gen_length} tokens in {trial_time:.3f}s "
|
||||
f"({gen_length / trial_time:.1f} tok/s)"
|
||||
@@ -199,6 +206,7 @@ class GenerationBenchmark:
|
||||
"gen_length": gen_length,
|
||||
"dtype": str(self.dtype),
|
||||
"device": self.device,
|
||||
"cache": self.cache_type,
|
||||
},
|
||||
)
|
||||
|
||||
@@ -216,6 +224,42 @@ def print_benchmark_result(result: BenchmarkResult):
|
||||
|
||||
|
||||
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(
|
||||
vocab_size=10000,
|
||||
dim=1536,
|
||||
@@ -227,23 +271,29 @@ if __name__ == "__main__":
|
||||
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("Running AutoRegressiveLM Generation Benchmark (KVCache)")
|
||||
print(
|
||||
f"Running AutoRegressiveLM Benchmark (device={args.device}, dtype={args.dtype})"
|
||||
)
|
||||
print("=" * 80)
|
||||
|
||||
prefill_result = benchmark.run_prefill_benchmark(
|
||||
batch_size=4,
|
||||
prompt_length=512,
|
||||
num_trials=5,
|
||||
)
|
||||
print_benchmark_result(prefill_result)
|
||||
if not args.decode_only:
|
||||
prefill_result = benchmark.run_prefill_benchmark(
|
||||
batch_size=args.batch_size,
|
||||
prompt_length=args.prompt_length,
|
||||
num_trials=args.num_trials,
|
||||
)
|
||||
print_benchmark_result(prefill_result)
|
||||
|
||||
gen_result = benchmark.run_decoding_benchmark(
|
||||
batch_size=4,
|
||||
prompt_length=512,
|
||||
gen_length=128,
|
||||
num_trials=5,
|
||||
)
|
||||
print_benchmark_result(gen_result)
|
||||
if not args.prefill_only:
|
||||
gen_result = benchmark.run_decoding_benchmark(
|
||||
batch_size=args.batch_size,
|
||||
prompt_length=args.prompt_length,
|
||||
gen_length=args.gen_length,
|
||||
num_trials=args.num_trials,
|
||||
)
|
||||
print_benchmark_result(gen_result)
|
||||
|
||||
+112
-32
@@ -1,7 +1,10 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from astrai.inference import InferenceEngine
|
||||
from astrai.model import AutoModel
|
||||
@@ -17,62 +20,109 @@ def processor(
|
||||
top_p: float,
|
||||
question_key: str,
|
||||
response_key: str,
|
||||
max_tokens: int,
|
||||
max_tokens: Optional[int],
|
||||
batch_size: int,
|
||||
num_samples: int = 1,
|
||||
cache_len: int = 2048,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
):
|
||||
# Load model and tokenizer
|
||||
print(f"Loading model from {param_path} ...")
|
||||
t0 = time.time()
|
||||
model = AutoModel.from_pretrained(param_path)
|
||||
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
||||
model.to(device="cuda", dtype=torch.bfloat16)
|
||||
print(f" model loaded in {time.time() - t0:.1f}s")
|
||||
|
||||
# Create inference engine
|
||||
engine = InferenceEngine(
|
||||
model=model, tokenizer=tokenizer, max_batch_size=batch_size
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=batch_size * num_samples,
|
||||
max_seq_len=cache_len,
|
||||
max_prompt_len=cache_len,
|
||||
)
|
||||
|
||||
print(f"Reading {input_json_file} ...")
|
||||
with open(input_json_file, "r", encoding="utf-8") as f:
|
||||
input_data = [json.loads(line) for line in f]
|
||||
|
||||
# Check input format: chat messages or raw text
|
||||
if input_data and "messages" in input_data[0]:
|
||||
# Chat format: [{"messages": [...]}]
|
||||
prompts = [
|
||||
tokenizer.apply_chat_template(item["messages"], tokenize=False)
|
||||
for item in input_data
|
||||
]
|
||||
else:
|
||||
# Raw text format: [{"question": "..."}]
|
||||
prompts = [item[question_key] for item in input_data]
|
||||
print(f" {len(prompts)} prompts loaded\n")
|
||||
|
||||
# Use provided max_tokens or default to model config max_len
|
||||
if max_tokens is None:
|
||||
max_tokens = model.config.max_len
|
||||
|
||||
# Generate responses (batch)
|
||||
responses = engine.generate(
|
||||
prompt=prompts,
|
||||
stream=False,
|
||||
max_tokens=max_tokens,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
)
|
||||
chunk_size = max(1, batch_size)
|
||||
|
||||
# Write results
|
||||
with open(output_json_file, "w", encoding="utf-8") as f:
|
||||
for prompt, response in zip(prompts, responses):
|
||||
if input_data and "messages" in input_data[0]:
|
||||
output_item = {"response": response}
|
||||
pbar = tqdm(
|
||||
total=len(prompts) * num_samples,
|
||||
unit="gen",
|
||||
desc=f" Generating ({num_samples}x/prompt)",
|
||||
)
|
||||
for chunk_start in range(0, len(prompts), chunk_size):
|
||||
chunk = prompts[chunk_start : chunk_start + chunk_size]
|
||||
|
||||
if num_samples > 1:
|
||||
chunk_expanded = [p for p in chunk for _ in range(num_samples)]
|
||||
resp_chunk = engine.generate(
|
||||
prompt=chunk_expanded,
|
||||
stream=False,
|
||||
max_tokens=max_tokens,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
frequency_penalty=frequency_penalty,
|
||||
rep_window=rep_window,
|
||||
)
|
||||
resp_chunk = [
|
||||
resp_chunk[i * num_samples : (i + 1) * num_samples]
|
||||
for i in range(len(chunk))
|
||||
]
|
||||
else:
|
||||
output_item = {question_key: prompt, response_key: response}
|
||||
f.write(json.dumps(output_item, ensure_ascii=False) + "\n")
|
||||
resp_chunk = engine.generate(
|
||||
prompt=chunk,
|
||||
stream=False,
|
||||
max_tokens=max_tokens,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
frequency_penalty=frequency_penalty,
|
||||
rep_window=rep_window,
|
||||
)
|
||||
|
||||
for i, prompt in enumerate(chunk):
|
||||
if input_data and "messages" in input_data[0]:
|
||||
orig = input_data[chunk_start + i]
|
||||
output_item = {**orig, response_key: resp_chunk[i]}
|
||||
else:
|
||||
output_item = {
|
||||
question_key: prompt,
|
||||
response_key: resp_chunk[i],
|
||||
}
|
||||
f.write(json.dumps(output_item, ensure_ascii=False) + "\n")
|
||||
|
||||
pbar.update(len(chunk) * num_samples)
|
||||
|
||||
pbar.close()
|
||||
|
||||
elapsed = time.time() - t0
|
||||
print(
|
||||
f"\nDone! {len(prompts)} prompts x {num_samples} samples -> {output_json_file}"
|
||||
)
|
||||
print(f"Total time: {elapsed:.1f}s ({elapsed / len(prompts):.2f}s/prompt)")
|
||||
|
||||
# Cleanup
|
||||
engine.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="Run generate with a Khaosz model.")
|
||||
parser = argparse.ArgumentParser(description="Batch generation from JSONL file.")
|
||||
|
||||
parser.add_argument(
|
||||
"--param_path", type=str, required=True, help="Path to the model directory."
|
||||
@@ -93,38 +143,68 @@ if __name__ == "__main__":
|
||||
"--question_key",
|
||||
type=str,
|
||||
default="question",
|
||||
help="Key for the question in the input JSON.",
|
||||
help="Key for the question in the input JSON (default: question).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--response_key",
|
||||
type=str,
|
||||
default="response",
|
||||
help="Key for the response in the output JSON.",
|
||||
help="Key for the response in the output JSON (default: response).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--temperature",
|
||||
type=float,
|
||||
default=0.60,
|
||||
help="Temperature for generating responses.",
|
||||
help="Temperature for generating responses (default: 0.60).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--top_k", type=int, default=30, help="Top-k value for generating responses."
|
||||
"--top_k",
|
||||
type=int,
|
||||
default=30,
|
||||
help="Top-k value for generating responses (default: 30).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--top_p",
|
||||
type=float,
|
||||
default=0.95,
|
||||
help="Top-p value for generating responses.",
|
||||
help="Top-p value for generating responses (default: 0.95).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size", type=int, default=1, help="Batch size for generating responses."
|
||||
"--batch_size",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Batch size for generating responses (default: 1).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num_samples",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of responses per prompt (expands batch internally, default: 1).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_tokens",
|
||||
type=int,
|
||||
default=2048,
|
||||
default=None,
|
||||
help="Maximum tokens to generate (default: model config max_len).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--cache_len",
|
||||
type=int,
|
||||
default=2048,
|
||||
help="KV cache & prompt truncation length (default: 2048, lower = less memory).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--frequency_penalty",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="Frequency penalty to reduce repetition (default: 0.0, try 0.5-1.0).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--rep_window",
|
||||
type=int,
|
||||
default=64,
|
||||
help="Window size for frequency penalty (default: 64).",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
|
||||
+211
-60
@@ -1,17 +1,98 @@
|
||||
import argparse
|
||||
import os
|
||||
from functools import partial
|
||||
from typing import Any, Dict
|
||||
|
||||
import torch
|
||||
import torch.optim as optim
|
||||
from torch import Tensor, nn
|
||||
|
||||
from astrai.config import AutoRegressiveLMConfig, TrainConfig
|
||||
from astrai.dataset import DatasetFactory
|
||||
from astrai.dataset import DatasetFactory, dpo_collate_fn, grpo_collate_fn
|
||||
from astrai.model import AutoRegressiveLM
|
||||
from astrai.model.components.decoder_block import DecoderBlock
|
||||
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"])
|
||||
self.param_groups = [*self.muon.param_groups, *self.adamw.param_groups]
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
|
||||
parser = argparse.ArgumentParser(description="Train the AutoRegressiveLM model.")
|
||||
@@ -35,6 +116,13 @@ def parse_args() -> argparse.Namespace:
|
||||
required=True,
|
||||
help="Path to the model parameters or resume checkpoint.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--resume",
|
||||
action="store_true",
|
||||
default=False,
|
||||
help="Resume training from checkpoint at --param_path "
|
||||
"(restore epoch, consumed_samples, optimizer & scheduler state).",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--n_epoch", type=int, default=1, help="Number of epochs to train."
|
||||
@@ -60,26 +148,39 @@ def parse_args() -> argparse.Namespace:
|
||||
parser.add_argument(
|
||||
"--max_grad_norm",
|
||||
type=float,
|
||||
default=1.0,
|
||||
help="Max gradient norm for clipping.",
|
||||
default=None,
|
||||
help="Max gradient norm for clipping. None disables clipping.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--adamw_beta1",
|
||||
"--weight_decay",
|
||||
type=float,
|
||||
default=0.9,
|
||||
help="Beta1 for AdamW optimizer.",
|
||||
default=0.1,
|
||||
help="Weight decay (applied to Muon matrix params; non-matrix use 0).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--adamw_beta2",
|
||||
"--muon_momentum",
|
||||
type=float,
|
||||
default=0.95,
|
||||
help="Beta2 for AdamW optimizer.",
|
||||
help="Momentum factor for Muon optimizer.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--adamw_weight_decay",
|
||||
type=float,
|
||||
default=0.01,
|
||||
help="Weight decay for AdamW optimizer.",
|
||||
"--muon_nesterov",
|
||||
action=argparse.BooleanOptionalAction,
|
||||
default=True,
|
||||
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(
|
||||
"--random_seed", type=int, default=3407, help="Random seed for reproducibility."
|
||||
@@ -150,8 +251,8 @@ def parse_args() -> argparse.Namespace:
|
||||
parser.add_argument(
|
||||
"--metrics",
|
||||
nargs="*",
|
||||
default=["loss", "lr"],
|
||||
help="Metrics to log (e.g. --metrics loss lr val_loss). Default: loss lr.",
|
||||
default=["loss", "lr", "grad_norm"],
|
||||
help="Metrics to log (e.g. --metrics loss lr val_loss). Default: loss lr grad_norm.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--log_dir",
|
||||
@@ -159,23 +260,14 @@ def parse_args() -> argparse.Namespace:
|
||||
default="checkpoint/logs",
|
||||
help="Directory for metric logs.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--log_interval",
|
||||
type=int,
|
||||
default=100,
|
||||
help="Number of batch iterations between metric logs.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--grpo_sync_interval",
|
||||
type=int,
|
||||
default=200,
|
||||
help="GRPO ref model sync interval (steps).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--start_epoch", type=int, default=0, help="Start epoch for training."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--start_batch", type=int, default=0, help="Start batch for training."
|
||||
"--start_samples",
|
||||
type=int,
|
||||
default=0,
|
||||
help="Start samples (per rank) for training.",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
@@ -221,6 +313,44 @@ def parse_args() -> argparse.Namespace:
|
||||
help="NEFTune noise alpha (0=disabled, typical: 5.0).",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--schedule_type",
|
||||
type=str,
|
||||
default="cosine",
|
||||
choices=["cosine", "sgdr", "wsd"],
|
||||
help="Learning rate scheduler type.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--min_rate",
|
||||
type=float,
|
||||
default=None,
|
||||
help="Minimum LR as fraction of base LR. Uses scheduler default if not set (cosine/sgdr: 0.05, wsd: 0.0).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--cycle_length",
|
||||
type=int,
|
||||
default=None,
|
||||
help="SGDR first cycle length in steps. Defaults to total_steps - warmup_steps.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--t_mult",
|
||||
type=int,
|
||||
default=2,
|
||||
help="SGDR cycle length multiplier per restart.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--stable_steps",
|
||||
type=int,
|
||||
default=None,
|
||||
help="WSD stable plateau steps. Required when --schedule_type wsd.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--decay_steps",
|
||||
type=int,
|
||||
default=None,
|
||||
help="WSD decay steps. Defaults to total_steps - warmup_steps - stable_steps.",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
return args
|
||||
@@ -230,8 +360,8 @@ def create_model(config):
|
||||
return AutoRegressiveLM(config).to(dtype=torch.bfloat16)
|
||||
|
||||
|
||||
def create_optimizer(model, **kwargs) -> optim.Optimizer:
|
||||
return optim.AdamW(model.parameters(), fused=True, **kwargs)
|
||||
def create_optimizer(model, **kwargs) -> MuonMix:
|
||||
return MuonMix(model, **kwargs)
|
||||
|
||||
|
||||
def create_scheduler(
|
||||
@@ -262,11 +392,11 @@ def train(
|
||||
train_type: str,
|
||||
param_path: str,
|
||||
data_root_path: str,
|
||||
max_lr: float,
|
||||
resume: bool,
|
||||
n_epoch: int,
|
||||
batch_per_device: int,
|
||||
start_epoch: int,
|
||||
start_batch: int,
|
||||
start_samples: int,
|
||||
grad_accum_steps: int,
|
||||
warmup_ratio: float,
|
||||
ckpt_interval: int,
|
||||
@@ -275,17 +405,7 @@ def train(
|
||||
val_step: int,
|
||||
metrics: list[str],
|
||||
log_dir: str,
|
||||
log_interval: int,
|
||||
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,
|
||||
label_smoothing: float,
|
||||
random_seed: int,
|
||||
num_workers: int,
|
||||
pin_memory: bool,
|
||||
@@ -300,6 +420,13 @@ def train(
|
||||
master_port: str,
|
||||
start_method: str,
|
||||
neftune_alpha: float,
|
||||
schedule_type: str,
|
||||
min_rate: float,
|
||||
cycle_length: int,
|
||||
t_mult: int,
|
||||
stable_steps: int,
|
||||
decay_steps: int,
|
||||
**kwargs,
|
||||
):
|
||||
assert train_type in ["seq", "sft", "dpo", "grpo"]
|
||||
assert os.path.exists(param_path)
|
||||
@@ -309,17 +436,17 @@ def train(
|
||||
# Load config
|
||||
config_path = os.path.join(param_path, "config.json")
|
||||
config = AutoRegressiveLMConfig.from_file(config_path)
|
||||
config.neftune_alpha = neftune_alpha
|
||||
|
||||
if window_size is None:
|
||||
window_size = config.max_len
|
||||
|
||||
strategy_kwargs = {
|
||||
"beta": dpo_beta,
|
||||
"label_smoothing": label_smoothing,
|
||||
"clip_eps": grpo_clip_eps,
|
||||
"kl_coef": grpo_kl_coef,
|
||||
"group_size": group_size,
|
||||
"sync_interval": grpo_sync_interval,
|
||||
"beta": kwargs.pop("dpo_beta"),
|
||||
"label_smoothing": kwargs.pop("label_smoothing"),
|
||||
"clip_eps": kwargs.pop("grpo_clip_eps"),
|
||||
"kl_coef": kwargs.pop("grpo_kl_coef"),
|
||||
"group_size": kwargs.pop("group_size"),
|
||||
}
|
||||
|
||||
executor_kwargs = {
|
||||
@@ -333,33 +460,57 @@ def train(
|
||||
load_path=data_root_path,
|
||||
window_size=window_size,
|
||||
stride=stride,
|
||||
tokenizer_path=param_path,
|
||||
)
|
||||
|
||||
optimizer_fn = partial(
|
||||
create_optimizer,
|
||||
**{
|
||||
"lr": max_lr,
|
||||
"betas": (adamw_beta1, adamw_beta2),
|
||||
"weight_decay": adamw_weight_decay,
|
||||
},
|
||||
lr=kwargs.pop("max_lr"),
|
||||
weight_decay=kwargs.pop("weight_decay"),
|
||||
momentum=kwargs.pop("muon_momentum"),
|
||||
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(
|
||||
len(dataset), n_epoch, batch_per_device, nprocs, grad_accum_steps
|
||||
)
|
||||
warmup_steps = int(warmup_ratio * total_steps)
|
||||
warmup_steps = min(warmup_steps, total_steps)
|
||||
|
||||
scheduler_kwargs = {"warmup_steps": warmup_steps}
|
||||
|
||||
if schedule_type == "cosine":
|
||||
scheduler_kwargs["lr_decay_steps"] = total_steps - warmup_steps
|
||||
elif schedule_type == "sgdr":
|
||||
scheduler_kwargs["cycle_length"] = cycle_length or (total_steps - warmup_steps)
|
||||
scheduler_kwargs["t_mult"] = t_mult
|
||||
elif schedule_type == "wsd":
|
||||
remaining = total_steps - warmup_steps
|
||||
stable_steps_ = stable_steps or max(1, int(remaining * 0.8))
|
||||
scheduler_kwargs["stable_steps"] = stable_steps_
|
||||
scheduler_kwargs["decay_steps"] = max(
|
||||
1, decay_steps or (remaining - stable_steps_)
|
||||
)
|
||||
|
||||
if min_rate is not None:
|
||||
scheduler_kwargs["min_rate"] = min_rate
|
||||
|
||||
scheduler_fn = partial(
|
||||
create_scheduler,
|
||||
**{
|
||||
"schedule_type": "cosine",
|
||||
"warmup_steps": min(warmup_steps, total_steps),
|
||||
"lr_decay_steps": total_steps - min(warmup_steps, total_steps),
|
||||
},
|
||||
schedule_type=schedule_type,
|
||||
**scheduler_kwargs,
|
||||
)
|
||||
|
||||
grad_ckpt_modules = [DecoderBlock] if gradient_checkpointing else []
|
||||
|
||||
collate_fn = None
|
||||
if train_type == "dpo":
|
||||
collate_fn = dpo_collate_fn
|
||||
elif train_type == "grpo":
|
||||
collate_fn = grpo_collate_fn
|
||||
|
||||
train_config = TrainConfig(
|
||||
model_fn=model_fn,
|
||||
strategy=train_type,
|
||||
@@ -370,7 +521,7 @@ def train(
|
||||
n_epoch=n_epoch,
|
||||
batch_per_device=batch_per_device,
|
||||
start_epoch=start_epoch,
|
||||
start_batch=start_batch,
|
||||
start_samples=start_samples,
|
||||
ckpt_interval=ckpt_interval,
|
||||
grad_accum_steps=grad_accum_steps,
|
||||
max_grad_norm=max_grad_norm,
|
||||
@@ -388,15 +539,15 @@ def train(
|
||||
val_step=val_step,
|
||||
metrics=metrics,
|
||||
log_dir=log_dir,
|
||||
log_interval=log_interval,
|
||||
gradient_checkpointing_modules=grad_ckpt_modules,
|
||||
executor_kwargs=executor_kwargs,
|
||||
extra_kwargs=strategy_kwargs,
|
||||
neftune_alpha=neftune_alpha,
|
||||
collate_fn=collate_fn,
|
||||
)
|
||||
|
||||
trainer = Trainer(train_config)
|
||||
trainer.train(resume_dir=param_path)
|
||||
trainer.train(param_path=param_path, resume=resume)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -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)
|
||||
+1
-1
@@ -75,7 +75,7 @@ class MultiTurnDataset(Dataset):
|
||||
|
||||
|
||||
class EarlyStoppingDataset(Dataset):
|
||||
"""Dataset that triggers early stopping after a specified number of iterations."""
|
||||
"""Dataset that triggers early stopping after consuming a specified number of samples."""
|
||||
|
||||
def __init__(self, length=10, stop_after=5):
|
||||
self.length = length
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
import json
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
import pytest
|
||||
@@ -8,6 +10,11 @@ from astrai.config.preprocess_config import (
|
||||
PipelineConfig,
|
||||
ProcessingConfig,
|
||||
)
|
||||
from astrai.preprocessing.builder import (
|
||||
MultiOutputMaskBuilder,
|
||||
SectionedMaskBuilder,
|
||||
SingleOutputMaskBuilder,
|
||||
)
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
_SPECIAL_TOKENS_CONFIG = {
|
||||
@@ -200,3 +207,43 @@ def make_grpo_no_template_config():
|
||||
mask_default="mask",
|
||||
preprocessing=ProcessingConfig(max_seq_len=2048),
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def builder():
|
||||
return SectionedMaskBuilder()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def single_builder():
|
||||
return SingleOutputMaskBuilder()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def multi_builder():
|
||||
return MultiOutputMaskBuilder()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def tokenizer_dir(temp_dir, test_tokenizer):
|
||||
d = os.path.join(temp_dir, "tok")
|
||||
os.makedirs(d, exist_ok=True)
|
||||
test_tokenizer._tokenizer.save(os.path.join(d, "tokenizer.json"))
|
||||
with open(os.path.join(d, "tokenizer_config.json"), "w") as f:
|
||||
json.dump(
|
||||
{"special_tokens": {"pad_token": "<|_pad_|>", "unk_token": "<|_unk_|>"}}, f
|
||||
)
|
||||
return d
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def chat_tokenizer_dir(temp_dir, chat_tokenizer):
|
||||
d = os.path.join(temp_dir, "tok")
|
||||
os.makedirs(d, exist_ok=True)
|
||||
chat_tokenizer._tokenizer.save(os.path.join(d, "tokenizer.json"))
|
||||
with open(os.path.join(d, "tokenizer_config.json"), "w") as f:
|
||||
json.dump(
|
||||
{"special_tokens": _SPECIAL_TOKENS_CONFIG, "chat_template": _CHAT_TEMPLATE},
|
||||
f,
|
||||
)
|
||||
return d
|
||||
|
||||
@@ -25,7 +25,9 @@ def test_single_process():
|
||||
|
||||
scheduler.step()
|
||||
|
||||
checkpoint = Checkpoint(state_dict=model.state_dict(), epoch=3, iteration=30)
|
||||
checkpoint = Checkpoint(
|
||||
state_dict=model.state_dict(), epoch=3, consumed_samples=120
|
||||
)
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
checkpoint.save(tmpdir)
|
||||
@@ -33,7 +35,7 @@ def test_single_process():
|
||||
loaded_checkpoint = Checkpoint.load(tmpdir)
|
||||
|
||||
assert loaded_checkpoint.epoch == 3
|
||||
assert loaded_checkpoint.iteration == 30
|
||||
assert loaded_checkpoint.consumed_samples == 120
|
||||
|
||||
|
||||
def test_checkpoint_with_extra():
|
||||
@@ -46,7 +48,10 @@ def test_checkpoint_with_extra():
|
||||
"scheduler": {"last_epoch": 5},
|
||||
}
|
||||
checkpoint = Checkpoint(
|
||||
state_dict=model.state_dict(), epoch=1, iteration=10, extra=extra
|
||||
state_dict=model.state_dict(),
|
||||
epoch=1,
|
||||
consumed_samples=40,
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
@@ -77,7 +82,7 @@ def simple_training():
|
||||
checkpoint = Checkpoint(
|
||||
state_dict=model.state_dict(),
|
||||
epoch=2,
|
||||
iteration=10,
|
||||
consumed_samples=40,
|
||||
)
|
||||
|
||||
rank = get_rank()
|
||||
|
||||
+816
-158
File diff suppressed because it is too large
Load Diff
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user