docs: restructure to docs/, add guides and developer docs
- Rename assets/ to docs/, split into guides/ and developer/ - Add get-started.md: installation + 5-step quickstart - Add guides/evaluation.md: 7 eval scripts with CLI args - Add guides/distributed.md: DDP/FSDP, gradient accumulation, NCCL - Add developer/internals.md: loss formulas, RoPE, KV cache math - Add developer/cuda_kernels.md: build system, benchmarks, file layout - Fix storage_format doc in preprocessing.md - Update cross-references in README.md, README-zh-CN.md, Dockerfile
This commit is contained in:
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,131 @@
|
||||
# CUDA Kernels
|
||||
|
||||
AstrAI includes optional custom CUDA attention kernels for decode and prefill. These are **not built by default** and are **not yet wired into the model or inference path** — they are standalone kernels with benchmarks and tests.
|
||||
|
||||
## Overview
|
||||
|
||||
| Kernel | File | Description |
|
||||
|--------|------|-------------|
|
||||
| `attn_decode` | `attn_decode.cu` | Basic GQA decode attention |
|
||||
| `attn_prefill` | `attn_prefill.cu` | Basic GQA prefill attention |
|
||||
| `attn_paged_decode` | `attn_paged_decode.cu` | Paged KV cache decode attention |
|
||||
|
||||
Additionally, optimized `.cuh` variants with tensor-core MMA (Matrix Multiply-Accumulate) exist:
|
||||
|
||||
| Variant | File | Optimization |
|
||||
|---------|------|--------------|
|
||||
| Split-KV MMA decode | `attn_decode_split_kv_mma.cuh` | Split KV across waraps + MMA (sm_80+) |
|
||||
| Split-Q MMA prefill | `attn_prefill_split_q_mma.cuh` | Split Q across waraps + MMA (sm_80+) |
|
||||
| Paged split-KV MMA decode | `attn_paged_decode_split_kv_mma.cuh` | Paged cache + split-KV + MMA |
|
||||
|
||||
## Build System
|
||||
|
||||
### Auto-detection
|
||||
|
||||
Kernels are built when **both** of these conditions are met:
|
||||
1. `nvcc` is available on `PATH`
|
||||
2. `torch.cuda.is_available()` returns `True`
|
||||
|
||||
Unless `CSRC_KERNELS=false` is set explicitly.
|
||||
|
||||
### Manual build
|
||||
|
||||
```bash
|
||||
# During install
|
||||
CSRC_KERNELS=true pip install -e . --no-build-isolation
|
||||
|
||||
# Rebuild after editing .cu/.cuh files
|
||||
CSRC_KERNELS=true python setup.py build_ext --inplace
|
||||
# Output: astrai/extension/*.so
|
||||
```
|
||||
|
||||
### Architecture flags
|
||||
|
||||
`csrc/build.py` auto-detects the GPU compute capability and generates the appropriate `nvcc` gencode flag:
|
||||
|
||||
- **sm_80+** (Ampere and later): enables tensor-core MMA path (`mma.sync.m16n8k16.bf16`)
|
||||
- **Below sm_80**: adds `-DASTRAI_NO_MMA` to disable the MMA path at compile time
|
||||
|
||||
### Build configuration
|
||||
|
||||
```
|
||||
NVCC_FLAGS = -O3 --expt-relaxed-constexpr --use_fast_math
|
||||
--ptxas-options=-O3,-v --extra-device-vectorization --threads=8
|
||||
```
|
||||
|
||||
The `REGISTRY` in `csrc/build.py` lists all registered kernels (currently 3). Each entry maps a kernel name to its source files and build flags.
|
||||
|
||||
## Python Wrappers
|
||||
|
||||
`astrai/extension/ops.py` provides Python wrappers for each compiled kernel. When the `.so` is not available, wrappers **fall back to `torch.nn.functional.scaled_dot_product_attention`** (SDPA).
|
||||
|
||||
Interface:
|
||||
```
|
||||
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"
|
||||
```
|
||||
|
||||
> **Note**: Wrappers are not yet called from `model/transformer.py` or `inference/`. The model uses PyTorch's built attention. Integration is future work.
|
||||
|
||||
## Standalone Testing
|
||||
|
||||
Each `csrc/tests/*.cu` file has the `nvcc` compile command in its header comment. Example:
|
||||
|
||||
```bash
|
||||
nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \
|
||||
--ptxas-options=-O3,-v --extra-device-vectorization \
|
||||
csrc/tests/attn_decode_test.cu -o /tmp/test && /tmp/test
|
||||
```
|
||||
|
||||
Test files:
|
||||
- `attn_decode_test.cu` — basic decode kernel
|
||||
- `attn_paged_decode_test.cu` — paged decode kernel
|
||||
- `attn_prefill_test.cu` — prefill kernel
|
||||
|
||||
## Benchmarks
|
||||
|
||||
Hardware: NVIDIA L20 (sm_89, 46 GB), CUDA 12.8, driver 570.86.
|
||||
|
||||
Reproduce:
|
||||
```bash
|
||||
nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \
|
||||
--ptxas-options=-O3,-v --extra-device-vectorization \
|
||||
csrc/tests/attn_<name>_test.cu -o /tmp/test && /tmp/test
|
||||
```
|
||||
|
||||
## Known Optimization Targets
|
||||
|
||||
- **Decode D=256**: spill eliminated (BC=16 + STAGES=2), but still 248 regs — further tiling could help.
|
||||
- **Prefill single-batch**: bandwidth low (52 GB/s at q=kv=2048) — likely compute-bound but near L20 bf16 ceiling (~94 TFLOP/s).
|
||||
- **Decode single-batch**: bandwidth low (309 GB/s at kv=512) — L20 HBM ~864 GB/s theoretical; small kv underutilizes SMs despite split-KV.
|
||||
|
||||
## File Layout
|
||||
|
||||
```
|
||||
csrc/
|
||||
├── build.py # Build system: REGISTRY, _arch_flags, nvcc flags
|
||||
├── kernels/
|
||||
│ ├── attn_common.h # Shared attention utilities
|
||||
│ ├── attn_decode.cu # Basic decode kernel (registered)
|
||||
│ ├── attn_prefill.cu # Basic prefill kernel (registered)
|
||||
│ ├── attn_paged_decode.cu # Paged decode kernel (registered)
|
||||
│ ├── attn_decode_split_kv.cuh # Split-KV variant
|
||||
│ ├── attn_decode_split_kv_mma.cuh # Split-KV + MMA variant
|
||||
│ ├── attn_prefill_split_q.cuh # Split-Q variant
|
||||
│ ├── attn_prefill_split_q_mma.cuh # Split-Q + MMA variant
|
||||
│ ├── attn_paged_decode_split_kv.cuh # Paged + split-KV variant
|
||||
│ ├── attn_paged_decode_split_kv_mma.cuh # Paged + split-KV + MMA variant
|
||||
│ ├── attn_dispatchers.cuh # Kernel dispatch macros
|
||||
│ ├── attn_entry_utils.cuh # Entry point helpers
|
||||
│ ├── attn_mma_utils.cuh # MMA utilities
|
||||
│ └── attn_warp_utils.cuh # Warp-level utilities
|
||||
└── tests/
|
||||
├── test_utils.cuh # Shared test utilities
|
||||
├── attn_decode_test.cu # Decode kernel test
|
||||
├── attn_paged_decode_test.cu # Paged decode test
|
||||
└── attn_prefill_test.cu # Prefill kernel test
|
||||
```
|
||||
|
||||
> Document Update Time: 2026-07-30
|
||||
@@ -0,0 +1,132 @@
|
||||
# Data Flow
|
||||
|
||||
This document describes the data pipeline: from raw text to model input tensors. For creating preprocessing configs, see [Preprocessing Guide](../guides/preprocessing.md).
|
||||
|
||||
## Contents
|
||||
|
||||
- [Overview](#overview)
|
||||
- [Data Preparation](#data-preparation) — tokenization, format detection, backends
|
||||
- [Data Keys by Training Type](#data-keys-by-training-type)
|
||||
- [Dataset Architecture](#dataset-architecture)
|
||||
- [Sampler](#sampler)
|
||||
- [DataLoader](#dataloader)
|
||||
|
||||
## Overview
|
||||
|
||||
```
|
||||
JSONL Lines → Pipeline (mask builder) → Tokenized Tensors
|
||||
↓
|
||||
.h5 or .bin storage
|
||||
↓
|
||||
Store.load()
|
||||
↓
|
||||
Store.fetch(begin, end, keys)
|
||||
↓
|
||||
BaseDataset.__getitem__(idx)
|
||||
↓
|
||||
Sampler → DataLoader → Training / Inference
|
||||
```
|
||||
|
||||
## Data Preparation
|
||||
|
||||
Raw text is tokenized via `AutoTokenizer.encode()` and saved as HDF5 (`.h5`) or binary (`.bin` + `meta.json`) files with keyed tensor groups.
|
||||
|
||||
### Tokenization
|
||||
|
||||
The `Pipeline` reads JSONL lines, applies the mask builder (see [Preprocessing](../guides/preprocessing.md)), and produces flat token sequences:
|
||||
|
||||
```python
|
||||
# Per JSONL line: messages → chat template → token IDs + loss mask
|
||||
tokens = tokenizer.encode(rendered_text) # List[int]
|
||||
loss_mask = [0, 0, 0, 1, 1, 1, 1, 1, 1] # 0=masked, 1=train
|
||||
# Stored as flat tensors, packed with other lines by packing strategy
|
||||
```
|
||||
|
||||
The output `meta.json` records the storage format, key names, dtype, total token count, and tensor shapes for each shard.
|
||||
|
||||
### Format Detection
|
||||
|
||||
`detect_format(load_path)` inspects the path:
|
||||
|
||||
- If `load_path` is a file: checks suffix — `.h5`/`.hdf5` → `"h5"`, `.jsonl` → `"jsonl"`, unknown suffix raises `ValueError`
|
||||
- If `load_path` is a directory: recursively globs for `*.h5`/`*.hdf5` files → `"h5"`, `*.bin` + `**/meta.json` → `"bin"`, or `*.jsonl` + `dataset_config.json` → `"jsonl"`
|
||||
|
||||
### Store Backends
|
||||
|
||||
Storage format is auto-detected by `detect_format()`; backends are dispatched via registry:
|
||||
|
||||
```
|
||||
StoreFactory.create("h5") → H5Store
|
||||
StoreFactory.create("bin") → MmapStore
|
||||
StoreFactory.create("jsonl") → JsonlStore
|
||||
```
|
||||
|
||||
All three inherit `Store` (base, owns `_data`/`_cum`/`_offsets`/`_normalize`) plus the `Streamable` and `Recordable` mixins, so every backend supports both `fetch(begin, end, keys)` (stream) and `fetch_record(index, keys)` (record) APIs.
|
||||
|
||||
**H5Store**: Reads HDF5 files. Tensors are loaded into host memory and normalized into segmented storage. `segments_are_records=True` — each `data_i` dataset is one record.
|
||||
|
||||
**MmapStore**: Memory-maps `.bin` files. OS page cache sharing is native — no explicit `share_memory_()` needed. Uses `torch.from_numpy(np.memmap(...))`. `segments_are_records=False` — bin segments are contiguous streams; record access is driven by `_offsets` (written when `save_bin(..., record_keys=...)` was used at preprocessing time).
|
||||
|
||||
**JsonlStore**: On-the-fly tokenization of raw JSONL files at load time. Requires a `dataset_config.json` alongside the `.jsonl` files following the same `PipelineConfig` schema with an additional `tokenizer_path` field. Two modes: eager (default, applies `TokenizeTransform` to all records at load) and lazy (`processor=fn` given, defers tokenisation to `fetch_record` — used by DPO/GRPO).
|
||||
|
||||
All backends normalise tensors into `Store._data[Dict[str, List[Tensor]]]` + `Store._cum[Dict[str, List[int]]]` (cumulative lengths for bisect-based stream indexing) + `Store._offsets[Dict[str, List[int]]]` (per-record offsets for record-mode indexing). Nested keys (GRPO `responses`/`masks` as `List[List[Tensor]]`) are stored as-is and excluded from both bookkeepings — they are only accessed record-by-record.
|
||||
|
||||
## Data Keys by Training Type
|
||||
|
||||
| Type | Storage Keys | Access Mode |
|
||||
|------|-------------|-------------|
|
||||
| `seq` | `sequence` (→ input_ids, target_ids via offset-by-1) | stream (`fetch`) |
|
||||
| `sft` | `sequence`, `loss_mask`, `position_ids` | stream (`fetch`) |
|
||||
| `dpo` | `chosen`, `rejected`, `chosen_mask`, `rejected_mask` | record (`fetch_record`) |
|
||||
| `grpo` | `prompts`, `responses`, `masks`, `rewards` | record (`fetch_record`) |
|
||||
|
||||
## Dataset Architecture
|
||||
|
||||
```
|
||||
DatasetFactory.load(
|
||||
train_type, load_path=None, window_size=0, 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)
|
||||
→ _normalize(raw) # base Store, shared by both backends
|
||||
→ Store._data[Dict[str, List[Tensor]]]
|
||||
+ _cum[Dict[str, List[int]]] (stream mode)
|
||||
+ _offsets[Dict[str, List[int]]] (record mode)
|
||||
|
||||
Stream datasets (SEQ/SFT):
|
||||
BaseDataset.__getitem__(idx)
|
||||
→ get_index(idx) → [begin, end)
|
||||
→ Store.fetch(begin, end, keys) → Tensor / Dict[str, Tensor]
|
||||
|
||||
Record datasets (DPO/GRPO via RecordDataset):
|
||||
RecordDataset.__getitem__(idx)
|
||||
→ Store.fetch_record(idx, keys) → Tensor / Dict[str, Tensor]
|
||||
```
|
||||
|
||||
Class hierarchy: `BaseDataset` ← `SEQDataset` / `SFTDataset` (stream); `BaseDataset` ← `RecordDataset` ← `DPODataset` / `GRPODataset` (record).
|
||||
|
||||
`window_size` = max input length, `stride` = step between consecutive samples (defaults to `window_size`, optional). Only meaningful for stream datasets — record datasets ignore both. `storage_type` defaults to `None` (auto-detect via `detect_format`).
|
||||
|
||||
`tokenizer_path` triggers lazy on-the-fly tokenisation for record datasets on raw JSONL (DPO builds a `dpo_processor`; SEQ/SFT/pre-tokenised backends ignore it). `store` (pre-built `Store`) bypasses `load_path`/`storage_type`/`tokenizer_path` entirely — the caller controls Store construction.
|
||||
|
||||
`Store.fetch(begin, end, keys)` (stream mode, on `Streamable`): accepts a single key (`str`) returning a `Tensor`, or a list of keys returning `Dict[str, Tensor]`. Internally uses `bisect` across multi-segment tensors. Raises `RuntimeError("Store not loaded")` if called before `load()`.
|
||||
|
||||
`Store.fetch_record(index, keys)` (record mode, on `Recordable`): same key API. Uses `_offsets[key]` when present (bin layout with per-record offsets), otherwise indexes `_data[key]` directly (H5/JSONL where each segment is one record).
|
||||
|
||||
## Sampler
|
||||
|
||||
`ResumableDistributedSampler` supports checkpoint-aware distributed sampling:
|
||||
|
||||
- Tracks `start_epoch` / `start_iter` for resume
|
||||
- Shuffle via `torch.Generator(seed + epoch)`
|
||||
- Per-replica index slicing for DDP
|
||||
|
||||
## DataLoader
|
||||
|
||||
Standard PyTorch `DataLoader` with configurable `batch_size`, `num_workers`, `pin_memory`, `prefetch_factor`. Sampler produces indices; dataloader fetches tensor batches via `__getitem__`.
|
||||
|
||||
> Document Update Time: 2026-07-19
|
||||
@@ -0,0 +1,208 @@
|
||||
# Internals
|
||||
|
||||
Mathematical foundations and internal algorithms for AstrAI's training, inference, and preprocessing pipelines. For practical usage guides, see [Training](../guides/training.md), [Inference](../guides/inference.md), and [Preprocessing](../guides/preprocessing.md).
|
||||
|
||||
## Contents
|
||||
|
||||
- [Autoregression & Causal Masking](#autoregression--causal-masking)
|
||||
- [Rotary Position Embedding (RoPE)](#rotary-position-embedding-rope)
|
||||
- [Training Loss Formulas](#training-loss-formulas)
|
||||
- [Training Loop Internals](#training-loop-internals)
|
||||
- [Callback Lifecycle](#callback-lifecycle)
|
||||
- [KV Cache Mathematics](#kv-cache-mathematics)
|
||||
- [Mask Algorithm Internals](#mask-algorithm-internals)
|
||||
- [Gradient Accumulation Mechanics](#gradient-accumulation-mechanics)
|
||||
|
||||
## Autoregression & Causal Masking
|
||||
|
||||
Given a token sequence, the model predicts the probability of the next token. Each generated token is appended to the input and fed back, repeating until an end-of-sequence token or max length.
|
||||
|
||||
```
|
||||
sequence : [[1, 2, 3, 4, 5, 6]]
|
||||
input_ids: [[1, 2, 3, 4, 5]]
|
||||
target_ids: [[2, 3, 4, 5, 6]]
|
||||
```
|
||||
|
||||
A lower-triangular causal mask prevents attending to future positions:
|
||||
|
||||
```
|
||||
[[0, -inf, -inf, -inf, -inf],
|
||||
[0, 0, -inf, -inf, -inf],
|
||||
[0, 0, 0, -inf, -inf],
|
||||
[0, 0, 0, 0, -inf],
|
||||
[0, 0, 0, 0, 0]]
|
||||
```
|
||||
|
||||
This ensures position $i$ can only attend to positions $\leq i$, which is essential for autoregressive generation.
|
||||
|
||||
## Rotary Position Embedding (RoPE)
|
||||
|
||||
RoPE embeds position into Q/K vectors via complex rotation:
|
||||
|
||||
$$ q_i = R_i W_q x_i, \quad k_j = R_j W_k x_j, \quad q_i^T k_j = x_i^T W_q^T R_{i-j} W_k x_j $$
|
||||
|
||||
The complex rotation `freqs_cis` is pre-computed once (`cos, sin` pairs per position). `apply_rotary_emb` multiplies Q/K as complex numbers. The key property is that the dot product $q_i^T k_j$ depends only on the relative position $i - j$, not the absolute positions.
|
||||
|
||||
**Critical for inference**: RoPE is applied **before** KV cache write, not after. If applied after caching, position encoding drift occurs because cached K/V would have stale rotation factors.
|
||||
|
||||
## Training Loss Formulas
|
||||
|
||||
### SEQ (Pre-training)
|
||||
|
||||
Next-token cross-entropy with optional label smoothing:
|
||||
|
||||
$$ L_{\text{PT}} = -\sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta) $$
|
||||
|
||||
### SFT (Supervised Fine-Tuning)
|
||||
|
||||
Masked cross-entropy (`ignore_index=-100`) over response tokens only:
|
||||
|
||||
$$ L_{\text{SFT}} = -\sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta) $$
|
||||
|
||||
Prompt tokens are masked out via `loss_mask`; only response tokens contribute to the loss.
|
||||
|
||||
### DPO (Direct Preference Optimization)
|
||||
|
||||
Frozen reference model, preference margin via log-ratio:
|
||||
|
||||
$$ L_{\text{DPO}} = -\mathbb{E}\left[\log\sigma\left(\beta\log\frac{\pi_\theta(y_w\mid x)}{\pi_{\text{ref}}(y_w\mid x)} - \beta\log\frac{\pi_\theta(y_l\mid x)}{\pi_{\text{ref}}(y_l\mid x)}\right)\right] $$
|
||||
|
||||
Parameters: `beta=0.1`, `reduction="sum"`.
|
||||
|
||||
### GRPO (Group Relative Policy Optimization)
|
||||
|
||||
Token-level PPO with group-normalized advantages:
|
||||
|
||||
$$ \text{Advantage}_i = \frac{r_i - \mu}{\sigma + \epsilon} $$
|
||||
|
||||
$$ L_{\text{GRPO}} = -\mathbb{E}_t\left[\min\left(\rho_t A,\; \text{clip}\left(\rho_t, 1-\epsilon, 1+\epsilon\right)A\right)\right] + \lambda \cdot \mathbb{E}_t\left[\frac{\pi_{\text{ref}}}{\pi_\theta} - \log\frac{\pi_{\text{ref}}}{\pi_\theta} - 1\right] $$
|
||||
|
||||
Where $\rho_t = \pi_\theta(a_t|s_t) / \pi_{\text{old}}(a_t|s_t)$ is the per-token importance sampling ratio. Advantages are derived from scalar per-response rewards, group-normalized, and broadcast across all response tokens. Only response tokens contribute to the loss.
|
||||
|
||||
Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`.
|
||||
|
||||
## Training Loop Internals
|
||||
|
||||
Two-level loop: **epoch** → **batch**. Optimizer step fires every `grad_accum_steps` batches.
|
||||
|
||||
```
|
||||
on_train_begin
|
||||
model.train()
|
||||
on_epoch_begin
|
||||
for batch in dataloader:
|
||||
on_batch_begin
|
||||
with executor.accumulate(model):
|
||||
loss = strategy.compute_loss(batch)
|
||||
context.loss = loss.item()
|
||||
stand_loss = loss / executor.grad_accum_steps
|
||||
executor.backward(stand_loss)
|
||||
context.consumed_samples += (
|
||||
context.config.batch_per_device * context.world_size
|
||||
)
|
||||
on_batch_end
|
||||
|
||||
if executor.sync_gradients:
|
||||
on_optimizer_step
|
||||
optimizer.step()
|
||||
optimizer.zero_grad()
|
||||
if scheduler:
|
||||
scheduler.step()
|
||||
on_epoch_end
|
||||
on_train_end
|
||||
```
|
||||
|
||||
The loss is divided by `grad_accum_steps` before `backward()`, so accumulated gradients sum to the correct mean.
|
||||
|
||||
## Callback Lifecycle
|
||||
|
||||
| Hook | Fires | Default callback |
|
||||
|------|-------|-----------------|
|
||||
| `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`, `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` (JSONL + validation, rank-0), `progress_bar` (tqdm), `gradient_clipping` (always registered; computes grad norm, clips only when `max_grad_norm` is not `None`).
|
||||
|
||||
## KV Cache Mathematics
|
||||
|
||||
At decode time, only the last query token matters. All previous K/V are cached to avoid recomputation:
|
||||
|
||||
$$ o_n = \sum_j \text{softmax}\left(\frac{q_n k_j}{\sqrt{d_k}}\right) v_j $$
|
||||
|
||||
The cache stores $k_j$ and $v_j$ for all previous positions. At each decode step, only $q_n$ (the current query) is computed fresh, and attention is computed against the cached K/V.
|
||||
|
||||
**RoPE ordering**: RoPE is applied to Q/K **before** writing to the KV cache. This is essential because:
|
||||
1. The cached K values already contain the rotation for their original positions.
|
||||
2. The new Q is rotated for its current position.
|
||||
3. The dot product $q_n^T k_j$ then correctly depends on $n - j$ (relative position).
|
||||
|
||||
If RoPE were applied after caching, the rotation factors would be inconsistent between cached and new tokens.
|
||||
|
||||
### Cache Implementations
|
||||
|
||||
- **ContiguousCache**: Each task gets a fixed slot of `[max_seq_len, num_key_value_heads, head_dim]`. Simple, efficient for small-to-medium batch sizes.
|
||||
- **PageCache**: Paged KV cache with prefix sharing. Uses `PagePool` (allocator + LRU + prefix matching) and `Storage` (page tensors). Enables sharing of common prompt prefixes across requests.
|
||||
|
||||
## Mask Algorithm Internals
|
||||
|
||||
### Template mode (`template: true`)
|
||||
|
||||
1. Prepend BOS token (masked)
|
||||
2. For each message in the field's array:
|
||||
1. Render through `chat_template` for that single message
|
||||
2. Encode rendered text
|
||||
3. Apply mask rule for the message's role
|
||||
|
||||
### Non-template mode
|
||||
|
||||
Encode the field value as text. Mask value is 1 (train) or 0 (mask) per the section's `action`.
|
||||
|
||||
### Text config detection
|
||||
|
||||
When no section uses `template` and all sections have `action: "train"`, the builder omits `loss_mask` from the output — all tokens are trained.
|
||||
|
||||
### Position ID strategies
|
||||
|
||||
| Mode | Behavior |
|
||||
|------|----------|
|
||||
| `none` | No position IDs generated |
|
||||
| `doc_reset` | Reset position to 0 at each document boundary in packed sequences |
|
||||
| `continuous` | Continuous position IDs across packed documents |
|
||||
|
||||
Default is `doc_reset`, which ensures each document in a packed bin starts from position 0, preventing position encoding drift between unrelated documents.
|
||||
|
||||
## Gradient Accumulation Mechanics
|
||||
|
||||
Three cooperating layers enable gradient accumulation:
|
||||
|
||||
1. **`GradientState`** — tracks the micro-step counter. Fires `sync_gradients=True` every `grad_accum_steps` micro-batches. The counter is incremented at the **start** of `accumulate()`, before the forward pass.
|
||||
|
||||
2. **`executor._no_sync(model)`** — suppresses gradient synchronization on non-sync micro-steps:
|
||||
- `NoneExecutor`: `nullcontext` (nothing to skip)
|
||||
- `DDPExecutor`: `model.no_sync()` (PyTorch's built-in — skips all-reduce of gradient buckets)
|
||||
- `FSDPExecutor`: `set_requires_gradient_sync(False, recurse=True)` on each `FSDPModule` (FSDP2's native mechanism)
|
||||
|
||||
3. **`AccumOptimizer` / `AccumScheduler`** — wrap the real optimizer/scheduler. `step()` and `zero_grad()` are gated on `sync_gradients` — they only forward to the inner optimizer when the sync flag is True.
|
||||
|
||||
The loss is divided by `grad_accum_steps` before `backward()`, so gradients sum to the correct mean across micro-steps. `consumed_samples` increments by `batch_per_device * world_size` every micro-batch.
|
||||
|
||||
### Effective batch size
|
||||
|
||||
$$ \text{Effective batch} = \text{nprocs} \times \text{batch\_per\_device} \times \text{grad\_accum\_steps} $$
|
||||
|
||||
### Total optimizer steps
|
||||
|
||||
```
|
||||
samples_per_replica = ceil(dataset_len / nprocs)
|
||||
batches_per_replica = ceil(samples_per_replica / batch_per_device)
|
||||
total_steps = (batches_per_replica // grad_accum_steps) * n_epoch
|
||||
```
|
||||
|
||||
This accounts for data-parallel sharding — each rank processes `1/nprocs` of the dataset.
|
||||
|
||||
> Document Update Time: 2026-07-30
|
||||
Reference in New Issue
Block a user