24 Commits
Author SHA1 Message Date
ViperEkura d08a92c7bd feat: add frequency penalty to inference sampling pipeline
- Add FrequencyPenaltyStrategy (logit -= penalty * count)
- Per-task rep_window for penalty history lookup
- Wire through engine, task, executor, API layer
- Add --frequency_penalty and --rep_window to stream_chat.py
- 9 unit tests for frequency penalty strategy
2026-07-17 21:28:31 +08:00
ViperEkura a1ea26d367 fix: rewrite GRPO data pipeline for offline record-level access
- process_list_field returns List[List[int]] preserving per-response boundaries
- GRPODataset rewritten to record-level __getitem__ (no windowing/stride)
- grpo_collate_fn pads variable-length responses into [B, G, R] tensors
- JsonlStore detects nested List[List[int]] and stores List[Tensor] per record
- Store._normalize skips nested-list keys from cumsum bookkeeping
- Pipeline._flush handles nested lists without cross-record flattening
- Export grpo_collate_fn from astrai.dataset
- 6 new GRPO tests + 2 updated builder tests, 114 total pass
2026-07-17 14:34:41 +08:00
ViperEkura c17aa0dc54 fix: eval script bugs and add missing features
- evaluate_mmlu: fix double few-shot injection (build_prompt no longer
  adds few-shot, apply_chat handles it once)
- evaluate_humaneval: fix pass@k k-filtering to be per-problem instead
  of using first problem's n globally; reuse ProcessPoolExecutor across
  problems; fix closure UnboundLocalError in test_one; handle None in
  report when k > n
- evaluate_ifd: remove dead code (score_plain/score_messages); add
  multi-file/directory input support with --input_path/--output_dir;
  add summary.json aggregation and --max_samples; add --dtype flag
- evaluate_ppl: add --device and --dtype flags (was hardcoded to cuda)
- evaluate_ifeval: fix docstring path (scripts/tools -> scripts/eval)
- analyze_weights: add --output JSON export; fix dead code filter
  ("_norm" not in r was always True)
2026-07-17 14:02:58 +08:00
ViperEkura b12b24eadc feat: rewrite evaluate_ppl with token-level loss and multi-file support
- Support multiple input files, glob patterns, and directory input
- Add --token_level flag: per-record token_ids + log_probs JSONL output
- Add --max_samples for random subsampling per file
- LossAccumulator: streaming mode (histogram-based percentiles, low memory) vs exact mode (full token list)
- Token type analysis (ascii/cjk/non_ascii/special) when token_level=True
- Fix token_ids/log_probs alignment (shift offset)
- Cache frozenset(stop_ids) outside loop for performance
- Aggregate stats: mean/median/ppl/p50/p90/p95/p99
- Summary JSON with all datasets in one file
2026-07-17 13:14:15 +08:00
ViperEkura cd14d53707 feat: implement bfd_split packing strategy
- BFDSplitPacking splits over-length sequences into chunks before BFD
- All keys (loss_mask, position_ids, ...) split in lockstep for alignment
- No tokens lost vs bfd which truncates over-length sequences
- Tests: token preservation, chunk alignment, short unchanged, vs bfd
2026-07-17 12:38:56 +08:00
ViperEkura e220413035 feat: support raw JSON files in dataset pipeline and JsonlStore
- detect_format now recognizes .json directories as jsonl store
- JsonlStore loads .json arrays and dicts alongside .jsonl
- tokenizer_path defaults to dataset dir when omitted
- Pipeline._iter_items handles .json files (arrays/single dict)
- Tests: detect_format, seq load, self-contained dataset dir
2026-07-17 12:20:03 +08:00
ViperEkura 84ed2327f5 feat: add --resume flag to decouple weight loading from training resumption
- Add --resume bool flag to train.py CLI
- --param_path always loads weights only by default
- --resume restores epoch, consumed_samples, optimizer & scheduler
- Checkpoint.load() now preserves full meta dict
- Update test_early_stopping to use new param_path/resume API
2026-07-16 14:23:23 +08:00
ViperEkura b14f301730 fix: init last_ckpt_step and last_log_flush_step from context.optimizer_step 2026-07-15 22:15:55 +08:00
ViperEkura 0654b4b916 refactor: template combine kernel, fix mask bug, unify dispatch
- Template combine kernel, share macros, extract entry_utils helpers
- Fix mask indexing (pass stride not pre-multiplied base)
- Remove !p.use_mask — MMA handles mask
2026-07-15 21:44:17 +08:00
ViperEkura 1f0be382ad refactor: extract load_q_mma_frags template, unify comment style
- Add load_q_mma_frags<KD>() shared template in attn_mma_utils.cuh
- Replace ~15 duplicated Q-load lines in 3 MMA kernels
- Unify section header comment style to // ---- Section ----
- Remove duplicate separator line in attn_mma_utils.cuh
2026-07-15 19:07:18 +08:00
ViperEkura bb175fda91 fix: resume optimizer LR, step display, and consumed_samples alignment 2026-07-15 08:59:52 +08:00
ViperEkura 13998da15a fix: uninitialized strides in decode test and wrong stride helper in paged test
- decode test main() missing set_default_strides → illegal memory access
- paged test used set_default_strides on PagedAttentionParams → compile error
2026-07-14 23:58:30 +08:00
ViperEkura 57729fd92d refactor: stride-based attn interface with layout and causal mask
- Replace is_causal + causal_offset with unified causal_offset (-1 = off, >=0 = first Q pos)
- Causal and mask can now coexist (was mutually exclusive)
- Add stride-based addressing for Q/KV/O (layout-agnostic, zero-copy)
- Add layout param ("bhld"/"blhd") parsed in Python, passed as int to C++
- Support 2D [batch, kv_len] and 3D [batch, q_len, kv_len] mask
- Vectorize paged KV gather in Python fallback (was per-token Python loop)
- Extract shared helpers: compute_num_splits, alloc_split_partials, dispatch_head_dim
- Unify paged_decode entry via attn_pack_paged_params
- Update mma_softmax_tile for 3D mask with per-row qrow indexing
2026-07-14 21:34:42 +08:00
ViperEkura 2c7a71a9c0 refactor: separate old policy and ref model in GRPO strategy
- Split single ref_model into old_model (importance sampling ratio) and ref_model (frozen KL regularizer)
- Move ref_model/old_model creation from strategy __init__ to TrainContextBuilder, pass as explicit parameters
- Remove periodic sync_ref_model + sync_interval; add sync_old_model for external rollout loop to call
- DPOStrategy also receives ref_model from builder
- Fix std to use unbiased=False (population std per GRPO paper)
- Remove redundant tests (test_grpo_kl_zero_at_init, test_grpo_no_sync_interval_param)
- Remove --grpo_sync_interval CLI arg
2026-07-14 20:03:45 +08:00
ViperEkura 3e0007fc91 docs : fix factory lists and MaskBuilderFactory docs
- Add MaskBuilderFactory, StoreWriterFactory, PackingStrategyFactory, PositionIdStrategyFactory to architecture design patterns
- Clarify MaskBuilderFactory three registered names (single, multi, sectioned) in preprocessing docs
2026-07-13 15:18:06 +08:00
ViperEkura b092316385 feat : add distributed checkpoint via executor checkpoint_context
- Add checkpoint_context context manager to BaseExecutor with entry/exit barrier
- Add _gather_state_dict hook overridden per executor (template method)
- DDPExecutor skips unwrap on non-rank-0 to avoid redundant state_dict gather
- FSDPExecutor uses rank0_only=True to reduce memory on non-writers
- Remove redundant rank-0 guard from Checkpoint.save and manual barrier from Callback
2026-07-13 12:27:09 +08:00
ViperEkura 9bcd696580 fix: token-level ratio and prompt masking in GRPO strategy
- Mask prompt tokens to 0 so their logprobs excluded from ratio/KL
- Switch to token-level ratio + PPO clipping via reduction='none'
- Slice response token logprobs from full sequence output
- Replace k3 KL estimator with non-negative k1 estimator
- Fix epsilon from finifo.eps (~1e-38) to 1e-8
- Remove unused 'reduction' param from GRPOStrategy.__init__
- Clarify offline batch semantics in docstring
- Add 11 unit tests for masking, advantage, KL, sync, clipping
- Sync training.md and architecture.md docs
2026-07-12 21:24:09 +08:00
ViperEkura 8f89c82d55 chore: bump version to 1.3.9 2026-07-12 21:04:20 +08:00
ViperEkura 21871197d7 refactor: use actual q_len from input, remove dead num_splits init 2026-07-12 19:48:17 +08:00
ViperEkura 4c35d36146 fix: auto-assign free port in spawn_parallel_fn to avoid EADDRINUSE 2026-07-12 19:21:56 +08:00
ViperEkura 9aca62c26c perf: remove per-element sentinel checks in softmax + fix stale comments
- Replace 4*NC8 per-element -FLT_MAX comparisons with 2 row-level pn guards; masked entries naturally underflow via expf(-FLT_MAX - nm) ≈ 0; pn only guards all-masked-row edge where nm == -FLT_MAX (exp(0)=1 not 0); ~1-3% speedup on prefill (verified via standalone CUDA bench); correctness verified: prefill 4/4, decode 3/3, paged decode 13/13
- Fix stale comments: remove false 'pre-scale Q' claim, correct occupancy numbers, remove phantom sQ from smem description
2026-07-12 15:45:08 +08:00
ViperEkura b5cdea98ad refactor: remove ineffective __launch_bounds__ from prefill kernel
- Remove MIN_BLOCKS template param and __launch_bounds__ attribute
- Profiling shows smem (not registers) is the occupancy bottleneck for D>=64, making the hint a no-op
- D=64 sees 2-4% speedup, D=128 unchanged (smem-capped at 1 block/SM)
- Update comment blocks in kernel header and both dispatch sites
- Verified correctness via standalone CUDA test (max_err ~1e-4)
2026-07-12 15:00:53 +08:00
ViperEkura 69fecaf387 perf: double-buffer KV pipeline and Q direct-to-register in decode
- Double-buffered KV (STAGES=2) for D<=128: next tile cp.async overlaps current tile MMA compute, hiding global load latency
- Q loaded directly from global into mma A-operand registers, removing sQ staging and prologue syncwarp
- Predicated cp.async unifies full and partial tile paths, eliminating scalar fallback branch
- STAGES=1 fallback for D=256 (double-buffer would exceed smem budget)
- Applied to both contiguous and paged decode MMA kernels
- ~1.27x average speedup on L20 (sm_89), zero precision loss
2026-07-12 14:14:54 +08:00
ViperEkura fd6d25ad86 refactor: extract bench_kernel and dispatch_by_head_dim into test_utils 2026-07-12 00:05:04 +08:00
58 changed files with 3097 additions and 989 deletions
+71
View File
@@ -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
+1 -2
View File
@@ -499,7 +499,6 @@ classDiagram
+float clip_eps
+float kl_coef
+int group_size
+str reduction
+int sync_interval
+compute_loss(batch) Tensor
+sync_ref_model()
@@ -1219,7 +1218,7 @@ classDiagram
| 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 |
+1 -1
View File
@@ -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
+9 -3
View File
@@ -122,17 +122,23 @@ Parameters: `beta=0.1`, `reduction="mean"`. Keys: `chosen`, `rejected`, `chosen_
### 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{ref}}(a_t|s_t)$ is the
per-token probability ratio and the expectations are over valid response tokens.
Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`, `sync_interval=200`.
Keys: `prompts`, `responses`, `masks`, `rewards`.
+1 -1
View File
@@ -1,4 +1,4 @@
__version__ = "1.3.8"
__version__ = "1.3.9"
__author__ = "ViperEkura"
from astrai.config import (
+2
View File
@@ -1,6 +1,7 @@
from astrai.dataset.dataset import (
BaseDataset,
DatasetFactory,
grpo_collate_fn,
)
from astrai.dataset.sampler import ResumableDistributedSampler
from astrai.dataset.storage import (
@@ -21,6 +22,7 @@ from astrai.serialization import (
__all__ = [
"BaseDataset",
"DatasetFactory",
"grpo_collate_fn",
"Store",
"StoreFactory",
"H5Store",
+116 -16
View File
@@ -15,6 +15,49 @@ from astrai.dataset.storage import (
from astrai.factory import BaseFactory
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,
}
class BaseDataset(Dataset, ABC):
"""Abstract base class for all dataset types.
@@ -250,28 +293,85 @@ class DPODataset(BaseDataset):
@DatasetFactory.register("grpo")
class GRPODataset(BaseDataset):
"""Dataset for Group Relative Policy Optimization training."""
"""Dataset for offline Group Relative Policy Optimization.
Unlike the window-based datasets (SEQ/SFT/DPO), GRPO data is
record-structured: each sample is one prompt with its group of
responses and scalar rewards. There is no windowing or stride —
every record is an independent training unit.
Expected storage layout (produced by JsonlStore or pre-tokenized):
- ``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
"""
def __init__(self, window_size: int = 0, stride: int = 0, **kwargs):
super().__init__(window_size=window_size, stride=stride or window_size)
self._records: List[dict] = []
@property
def required_keys(self) -> List[str]:
return ["prompts", "responses", "masks", "rewards"]
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
return self.storage.fetch(begin_idx, end_idx, key)
def load(self, load_path: str, storage_type: Optional[str] = None, **kwargs):
if storage_type is None:
storage_type = detect_format(load_path)
self.storage = StoreFactory.create(storage_type, **kwargs)
self._load_path = load_path
self.storage.load(load_path, **kwargs)
self._validate_keys()
self._build_records()
def _validate_keys(self):
actual_keys = set(self.storage.keys)
missing = [k for k in self.required_keys if k not in actual_keys]
if missing:
raise KeyError(
f"GRPODataset requires keys {self.required_keys}, "
f"but storage only has {sorted(actual_keys)}. Missing: {missing}"
)
def _build_records(self):
"""Unfold segmented storage into per-record lists.
``prompts`` is a flat list of 1-D tensors (one per record).
``responses`` / ``masks`` are nested lists (G tensors per record).
``rewards`` is a flat list of 1-D tensors (len G per record).
"""
prompt_segs = self.storage._data.get("prompts", [])
response_segs = self.storage._data.get("responses", [])
mask_segs = self.storage._data.get("masks", [])
reward_segs = self.storage._data.get("rewards", [])
n_records = len(prompt_segs)
self._records = []
for i in range(n_records):
self._records.append(
{
"prompts": prompt_segs[i],
"responses": response_segs[i] if i < len(response_segs) else [],
"masks": mask_segs[i] if i < len(mask_segs) else [],
"rewards": reward_segs[i]
if i < len(reward_segs)
else torch.tensor([]),
}
)
@property
def count(self) -> int:
return len(self._records)
def __len__(self) -> int:
return len(self._records)
def __getitem__(self, index: int) -> Dict[str, Tensor]:
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")
rec = self._records[index]
return {
"prompts": prompts,
"responses": responses,
"masks": masks,
"rewards": rewards,
"prompts": rec["prompts"].to(dtype=torch.long),
"responses": [r.to(dtype=torch.long) for r in rec["responses"]],
"masks": [m.to(dtype=torch.bool) for m in rec["masks"]],
"rewards": rec["rewards"].to(dtype=torch.float32),
}
+61 -28
View File
@@ -81,6 +81,11 @@ def detect_format(load_path: str) -> str:
]
if jsonl_files:
return "jsonl"
json_files = [
Path(p) for p in glob.glob(str(root / "**" / "*.json"), recursive=True)
]
if json_files:
return "jsonl"
raise FileNotFoundError(f"No supported data files found at {load_path}")
@@ -143,26 +148,37 @@ class Store(ABC):
return results[0] if len(results) == 1 else torch.cat(results, dim=0)
def _normalize(self, raw: Dict[str, List[Tensor]]):
def _normalize(self, raw: Dict[str, list]):
"""Register segments and pre-compute cumulative lengths.
Does NOT concatenate — segments are kept as-is to avoid OOM on
large datasets. Sets ``self._length`` to the minimum total
element count across all keys.
element count across all flat-tensor keys.
For GRPO multi-response keys, values may be ``List[List[Tensor]]``
(one list of G tensors per record). These are stored as-is and
excluded from the cumulative-length bookkeeping since they are
accessed record-by-record via ``_data`` rather than via ``fetch``.
"""
flat_lengths = []
for key, tensors in raw.items():
self._data[key] = tensors
if not tensors:
self._cum[key] = []
flat_lengths.append(0)
continue
# Skip nested lists (GRPO responses/masks) — record-level access
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
self._length = (
min((cum[-1] if cum else 0) for cum in self._cum.values())
if self._cum
else 0
)
flat_lengths.append(cum[-1] if cum else 0)
self._length = min(flat_lengths) if flat_lengths else 0
class StoreFactory(BaseFactory["Store"]):
@@ -245,12 +261,7 @@ class JsonlStore(Store):
with open(config_path, "r", encoding="utf-8") as f:
raw_config = json.load(f)
tokenizer_path = raw_config.pop("tokenizer_path", None)
if tokenizer_path is None:
raise ValueError(
f"JSONL dataset config must specify 'tokenizer_path': {config_path}"
)
tokenizer_path = raw_config.pop("tokenizer_path", None) or str(root)
self.config = PipelineConfig.from_dict(raw_config)
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
mask_builder = MaskBuilderFactory.create("sectioned")
@@ -261,6 +272,27 @@ class JsonlStore(Store):
raw: Dict[str, List[Tensor]] = {}
doc_sequences: List[List[int]] = []
def _process_item(item: dict) -> None:
nonlocal raw, doc_sequences
result = mask_builder.build(item, self.config, tokenizer)
if result is None:
return
result.pop("domain", None)
primary_ids = self._primary_ids(result)
if not primary_ids:
return
doc_sequences.append(primary_ids)
for key, ids in result.items():
if key not in raw:
raw[key] = []
if ids and isinstance(ids[0], list):
# GRPO multi-response: List[List[int]] → List[Tensor]
raw[key].append(
[torch.tensor(sub, dtype=self._infer_dtype(sub)) for sub in ids]
)
else:
raw[key].append(torch.tensor(ids, dtype=self._infer_dtype(ids)))
for jsonl_path in sorted(root.glob("*.jsonl")):
with open(jsonl_path, "r", encoding="utf-8") as f:
for line in f:
@@ -274,21 +306,22 @@ class JsonlStore(Store):
"Failed to parse JSON line in %s, skipping", jsonl_path
)
continue
_process_item(item)
result = mask_builder.build(item, self.config, tokenizer)
if result is None:
continue
result.pop("domain", None)
primary_ids = self._primary_ids(result)
if not primary_ids:
continue
doc_sequences.append(primary_ids)
for key, ids in result.items():
if key not in raw:
raw[key] = []
raw[key].append(torch.tensor(ids, dtype=self._infer_dtype(ids)))
for json_path in sorted(root.glob("*.json")):
if json_path.name == self.CONFIG_NAME:
continue
with open(json_path, "r", encoding="utf-8") as f:
try:
data = json.load(f)
except json.JSONDecodeError:
logger.warning("Failed to parse JSON file %s, skipping", json_path)
continue
if isinstance(data, list):
for item in data:
_process_item(item)
elif isinstance(data, dict):
_process_item(data)
pos_ids = position_strategy.generate(doc_sequences)
if pos_ids:
@@ -298,7 +331,7 @@ class JsonlStore(Store):
@staticmethod
def _primary_ids(result: dict) -> List[int]:
"""Return the first integer list in *result* as the primary id sequence."""
"""Return the first flat integer list in *result* as the primary id sequence."""
for val in result.values():
if isinstance(val, list) and val and isinstance(val[0], int):
return val
+8
View File
@@ -5,6 +5,14 @@ Public API:
- ``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``.
"""
+129 -32
View File
@@ -3,15 +3,43 @@
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
@@ -26,20 +54,84 @@ def _expand_kv_heads(
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,
is_causal: bool,
scale: float | None,
causal_offset: int,
scale: float,
q_layout: int,
kv_layout: int | None = None,
) -> torch.Tensor:
"""Reference attention via ``scaled_dot_product_attention``."""
"""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 = mask[:, None, None, :] if mask is not None else None
return F.scaled_dot_product_attention(
q, k, v, attn_mask=attn_mask, is_causal=is_causal and mask is None, scale=scale
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(
@@ -56,23 +148,22 @@ def _gather_kv_from_pages(
k_cache : [n_pages, page_size, n_kv_heads, head_dim]
v_cache : same as k_cache
Returns:
k, v : [batch, n_kv_heads, kv_len, head_dim]
k, v : [batch, kv_len, n_kv_heads, head_dim] (b l h d)
"""
batch, max_pages = page_table.shape
n_pages, ps, n_kv_heads, head_dim = k_cache.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}")
k = k_cache.new_empty(batch, n_kv_heads, kv_len, head_dim)
v = v_cache.new_empty(batch, n_kv_heads, kv_len, head_dim)
# 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]
for b in range(batch):
for pos in range(kv_len):
log_pg = pos // page_size
pg_off = pos % page_size
phys = int(page_table[b, log_pg].item())
k[b, :, pos, :] = k_cache[phys, pg_off, :, :]
v[b, :, pos, :] = v_cache[phys, pg_off, :, :]
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
@@ -81,21 +172,22 @@ def attn_decode(
k: torch.Tensor,
v: torch.Tensor,
mask: torch.Tensor | None = None,
is_causal: bool = False,
causal_offset: int = 0,
scale: float | 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,
is_causal=is_causal,
causal_offset=causal_offset,
scale=scale,
layout=li,
)
return _torch_fallback(q, k, v, mask, is_causal, scale)
return _torch_fallback(q, k, v, mask, causal_offset, scale, q_layout=li)
def attn_prefill(
@@ -103,21 +195,22 @@ def attn_prefill(
k: torch.Tensor,
v: torch.Tensor,
mask: torch.Tensor | None = None,
is_causal: bool = False,
causal_offset: int = 0,
scale: float | 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,
is_causal=is_causal,
causal_offset=causal_offset,
scale=scale,
layout=li,
)
return _torch_fallback(q, k, v, mask, is_causal, scale)
return _torch_fallback(q, k, v, mask, causal_offset, scale, q_layout=li)
def attn_paged_decode(
@@ -128,10 +221,11 @@ def attn_paged_decode(
page_size: int,
kv_len: int,
mask: torch.Tensor | None = None,
is_causal: bool = False,
causal_offset: int = 0,
scale: float | 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,
@@ -141,9 +235,12 @@ def attn_paged_decode(
page_size,
kv_len,
mask=mask,
is_causal=is_causal,
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, is_causal, scale)
return _torch_fallback(
q, k, v, mask, causal_offset, scale, q_layout=li, kv_layout=1
)
+3 -1
View File
@@ -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 (
@@ -50,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,
@@ -83,6 +84,7 @@ __all__ = [
"TemperatureStrategy",
"TopKStrategy",
"TopPStrategy",
"FrequencyPenaltyStrategy",
"SamplingPipeline",
"ProtocolHandler",
"StopChecker",
-1
View File
@@ -21,7 +21,6 @@ logger = logging.getLogger(__name__)
_UNSUPPORTED_PARAMS = (
"n",
"presence_penalty",
"frequency_penalty",
"logit_bias",
"user",
)
+1
View File
@@ -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:
+30
View File
@@ -75,6 +75,33 @@ class Executor:
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(
@@ -89,4 +116,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()
+8
View File
@@ -33,6 +33,8 @@ class Task:
temperature: float = 1.0,
top_p: float = 1.0,
top_k: int = 50,
frequency_penalty: float = 0.0,
rep_window: int = 64,
):
self.task_id = task_id
self.prompt_ids = prompt_ids
@@ -40,6 +42,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] = []
@@ -92,6 +96,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]}"
@@ -116,6 +122,8 @@ class TaskManager:
temperature=temperature,
top_p=top_p,
top_k=top_k,
frequency_penalty=frequency_penalty,
rep_window=rep_window,
)
with self._lock:
+63 -5
View File
@@ -74,6 +74,8 @@ 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):
@@ -82,12 +84,21 @@ class GenerationRequest:
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(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
@@ -132,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(
@@ -152,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():
@@ -185,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(
@@ -194,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)
@@ -206,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)
@@ -226,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
@@ -262,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:
+144 -13
View File
@@ -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,39 @@ 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:
@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).
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]``.
"""
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,6 +295,9 @@ 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).
@@ -180,6 +306,10 @@ def sample(
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]``.
@@ -189,5 +319,6 @@ def sample(
TemperatureStrategy(temperature),
TopKStrategy(top_k),
TopPStrategy(top_p),
FrequencyPenaltyStrategy(frequency_penalty),
]
).sample(logits, filter_value)
).sample(logits, filter_value, input_ids, input_mask)
+24 -1
View File
@@ -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
@@ -208,6 +224,13 @@ class DDPExecutor(BaseExecutor):
return model.module.state_dict()
return model.state_dict()
def _gather_state_dict(self, model: nn.Module):
if not self.use_distributed:
return self.unwrap_model(model)
if get_rank() != 0:
return None
return self.unwrap_model(model)
@ExecutorFactory.register("fsdp")
class FSDPExecutor(BaseExecutor):
@@ -279,7 +302,7 @@ class FSDPExecutor(BaseExecutor):
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()
+11 -2
View File
@@ -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"]
@@ -217,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(
+34 -21
View File
@@ -95,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"]
@@ -108,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)}
@@ -130,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:
@@ -282,10 +287,18 @@ class MultiOutputMaskBuilder(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
+47
View File
@@ -119,3 +119,50 @@ class BFDPacking(PackingStrategy):
bin_lengths.append(seq_len)
return bins
@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 _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
+30 -8
View File
@@ -149,11 +149,18 @@ class Pipeline:
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():
@@ -173,9 +180,24 @@ class Pipeline:
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)
]
# GRPO multi-response keys store List[List[int]] per record
# (responses/masks). Rewards store List[float] per record.
# Both produce List[Tensor] (one tensor per record), but
# responses need inner flattening while rewards do not.
if ids_list and isinstance(ids_list[0], list):
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)
]
if mode == "continuous" and original_sequences:
pos_ids = self._position_id.generate(keys.get("sequence", []))
+1 -4
View File
@@ -2,7 +2,6 @@
import io
import json
import os
import time
from dataclasses import dataclass, field
from pathlib import Path
@@ -148,9 +147,6 @@ class Checkpoint:
save_path = Path(save_dir)
save_path.mkdir(parents=True, exist_ok=True)
if get_rank() != 0:
return
meta = {
"epoch": self.epoch,
"consumed_samples": self.consumed_samples,
@@ -181,6 +177,7 @@ class Checkpoint:
epoch=meta.get("epoch", 0),
consumed_samples=meta.get("consumed_samples", 0),
extra=extra,
meta=meta,
config=config,
)
+65 -36
View File
@@ -223,14 +223,13 @@ class DPOStrategy(BaseStrategy):
self,
model: nn.Module,
device: str,
ref_model: nn.Module,
beta: float = 0.1,
reduction: str = "mean",
**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 +266,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 +315,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
+33 -23
View File
@@ -14,7 +14,7 @@ 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_norm,
@@ -139,28 +139,31 @@ 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_step = 0
self.last_ckpt_step = None
def _save_checkpoint(self, context: TrainContext):
state_dict = context.executor.unwrap_model(context.model)
def on_train_begin(self, context: TrainContext):
self.last_ckpt_step = context.optimizer_step
if get_rank() == 0:
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 _save_checkpoint(self, context: TrainContext):
self.last_ckpt_step = context.optimizer_step
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.optimizer_step - self.last_ckpt_step >= self.interval:
@@ -210,7 +213,7 @@ class ProgressBarCallback(TrainCallback):
@only_on_rank(0)
def on_optimizer_step(self, context: TrainContext):
postfix = {
"step": context.optimizer_step,
"step": f"{context.optimizer_step:d}",
"loss": f"{context.loss:.4f}",
"lr": f"{context.optimizer.param_groups[-1]['lr']:.2e}",
}
@@ -237,7 +240,7 @@ class MetricCallback(TrainCallback):
metrics: List[str] = None,
val_step: int = 0,
):
self.last_log_flush_step = 0
self.last_log_flush_step = None
self.save_interval = save_interval
self.metrics = metrics or ["loss", "lr"]
self.val_step = val_step
@@ -298,6 +301,9 @@ class MetricCallback(TrainCallback):
context.model.train()
return avg_loss
def on_train_begin(self, context: TrainContext):
self.last_log_flush_step = context.optimizer_step
@only_on_rank(0)
def _flush(self, epoch, step):
log_file = self.log_dir / f"epoch_{epoch}_step_{step}_metric.jsonl"
@@ -327,8 +333,12 @@ class MetricCallback(TrainCallback):
self._append("epoch", context)
def on_train_end(self, context):
if context.optimizer_step != self.last_log_flush_step:
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)
+42 -15
View File
@@ -13,7 +13,7 @@ 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
@@ -54,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:
@@ -74,8 +76,8 @@ class TrainContextBuilder:
model = model.to(device=device)
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)
@@ -91,18 +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
if checkpoint.consumed_samples > 0:
context.consumed_samples = checkpoint.consumed_samples
else:
context.consumed_samples = cfg.start_samples * context.world_size
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(
@@ -177,13 +190,27 @@ 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
+7 -4
View File
@@ -52,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)
@@ -95,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,
@@ -105,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,
)
+1 -1
View File
@@ -20,7 +20,7 @@ def _arch_flags() -> list[str]:
_kernels_dir = Path("csrc/kernels")
REGISTRY: dict[str, dict] = {}
CXX_FLAGS = ["-O3", "-march=native", "-funroll-loops"]
CXX_FLAGS = ["-O3", "-funroll-loops"]
NVCC_FLAGS = [
"-O3",
"--expt-relaxed-constexpr",
+18 -4
View File
@@ -10,11 +10,19 @@ struct AttentionParams {
int kv_len;
int head_dim;
int use_mask;
int is_causal;
int causal_offset;
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;
@@ -34,7 +42,6 @@ struct PagedAttentionParams {
int kv_len;
int head_dim;
int use_mask;
int is_causal;
int causal_offset;
float scale;
@@ -42,12 +49,19 @@ struct PagedAttentionParams {
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;
+23 -52
View File
@@ -5,24 +5,12 @@
#include "attn_decode_split_kv_mma.cuh"
#endif
static int decode_num_splits(int base_blocks, int tiles_total) {
int sm_count = 0;
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
return std::max(1, std::min(n, std::min(tiles_total, 32)));
}
// Scalar fallback: one warp per query head, split-KV across grid.z.
static void launch_scalar_decode(AttentionParams<bf16>& p) {
int group_size = p.q_head / p.kv_head;
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
p.num_splits = decode_num_splits(p.batch * p.kv_head, chunks_total);
auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
auto o_part = torch::empty({p.batch, p.q_head, p.num_splits, p.head_dim}, fopt);
auto ml_part = torch::empty({p.batch, p.q_head, p.num_splits, 2}, fopt);
p.o_part = o_part.data_ptr<float>();
p.ml_part = ml_part.data_ptr<float>();
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);
@@ -30,21 +18,18 @@ static void launch_scalar_decode(AttentionParams<bf16>& p) {
}
#ifndef ASTRAI_NO_MMA
// MMA head-packing requires G <= 16 (sQ has BR=16 rows). sm_80+ tensor-core
// 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).
template <int HEAD_DIM, int BC>
// STAGES=2 (double-buffer) for D<=128 (smem 16 KB); STAGES=1 for D=256
// (double-buffer would be 32 KB, near the 48 KB static cap — keep single
// to preserve occupancy).
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
static void launch_mma_decode(AttentionParams<bf16>& p) {
int tiles_total = (p.kv_len + BC - 1) / BC;
p.num_splits = decode_num_splits(p.batch * p.kv_head, tiles_total);
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
alloc_split_partials(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 = o_part.data_ptr<float>();
p.ml_part = ml_part.data_ptr<float>();
attn_decode_split_kv_mma_kernel<HEAD_DIM, BC>
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(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
@@ -53,7 +38,7 @@ template <int HEAD_DIM>
static void dispatch_decode(AttentionParams<bf16>& p) {
#ifndef ASTRAI_NO_MMA
int G = p.q_head / p.kv_head;
if (!p.use_mask && G >= 1 && G <= 16) {
if (G >= 1 && G <= 16) {
launch_mma_decode<HEAD_DIM, 32>(p);
return;
}
@@ -66,35 +51,21 @@ torch::Tensor attn_decode(
torch::Tensor k,
torch::Tensor v,
c10::optional<torch::Tensor> mask,
bool is_causal = false,
int64_t causal_offset = 0,
c10::optional<double> scale = c10::nullopt
int64_t causal_offset,
double scale,
int64_t layout
) {
AttentionParams<bf16> p;
attn_pack_params(q, k, v, mask, is_causal, causal_offset, scale, 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");
auto O = torch::empty_like(q);
p.o = (bf16*)O.data_ptr();
// 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();
switch (p.head_dim) {
case 32:
dispatch_decode<32>(p);
break;
case 64:
dispatch_decode<64>(p);
break;
case 128:
dispatch_decode<128>(p);
break;
case 256:
dispatch_decode<256>(p);
break;
default:
TORCH_CHECK(false, "decode: unsupported head_dim ", p.head_dim,
" (supported: 32, 64, 128, 256)");
}
DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p);
return O;
}
@@ -104,8 +75,8 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
py::arg("k"),
py::arg("v"),
py::arg("mask") = py::none(),
py::arg("is_causal") = false,
py::arg("causal_offset") = 0,
py::arg("scale") = py::none(),
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)");
}
+27 -11
View File
@@ -21,13 +21,16 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
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_head + q_head) * 1) * p.head_dim + lane * hd_per_thread;
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]);
q_reg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
int kv_base = ((batch * p.kv_head + kv_head) * p.kv_len) * p.head_dim;
int mask_base = batch * p.kv_len;
// 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};
@@ -43,9 +46,15 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
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)
k_smem[i] = p.k[kv_base + chunk_start * p.head_dim + i];
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++) {
@@ -54,9 +63,10 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
partial += q_reg[i] * __bfloat162float(k_smem[s * p.head_dim + lane * hd_per_thread + i]);
partial = warp_reduce_sum(partial) * p.scale;
if (p.use_mask && p.mask && !p.mask[mask_base + chunk_start + s])
int kv_idx = chunk_start + s;
if (p.use_mask && p.mask && !p.mask[mask_base + kv_idx])
partial = -FLT_MAX;
if (p.is_causal && (chunk_start + s) > p.causal_offset)
if (p.causal_offset >= 0 && kv_idx > p.causal_offset)
partial = -FLT_MAX;
float new_m = fmaxf(m, partial);
@@ -64,9 +74,10 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
float beta = expf(partial - new_m);
d = d * alpha + beta;
int v_off = kv_base + (chunk_start + s) * p.head_dim + lane * hd_per_thread;
// 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]) * beta;
acc_reg[i] = acc_reg[i] * alpha + __bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta;
m = new_m;
}
__syncthreads();
@@ -94,6 +105,9 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
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;
@@ -112,5 +126,7 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
}
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
p.o[(size_t)bh * p.head_dim + d] = __float2bfloat16(acc * inv);
// 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);
}
+80 -76
View File
@@ -8,36 +8,29 @@ 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,
// Decode has q_len == 1, so S = q @ K^T is a GEMV per head — no tensor-core
// work on its own. But GQA gives us G = q_head / kv_head query heads that all
// share one kv_head. We pack those G heads into the M=16 rows of
// mma.sync.m16n8k16, turning G independent GEMVs into a single GEMM that
// reuses each loaded K/V tile across all G heads (K/V load is the decode
// bottleneck, so the reuse is the win, not the flops). The KV sequence is
// partitioned across gridDim.z blocks so that a decode with only
// batch*kv_head independent tasks can fill all SMs. Each (batch, kv_head,
// split) block computes an UN-normalised partial (Oacc, m, l) over its KV
// slice; the combine kernel below reduces across splits. Fixes the "grid too
// small" bottleneck (0.04 waves/SM → many blocks) for long-context,
// small-batch decode.
//
// Partial layout (float, contiguous):
// o_part : [batch, q_head, num_splits, HEAD_DIM]
// ml_part: [batch, q_head, num_splits, 2] (m, l)
//
// Optimizations:
// - cp.async global→shared for K/V (bypasses registers, cuts instruction count)
// - XOR swizzle (swiz_col): LD=HEAD_DIM, zero waste, no bank conflicts
// - pre-scaled Q: Q scaled during load, softmax skips per-tile multiply
// - single-buffer: keeps smem small for high occupancy
template <int HEAD_DIM, int BC>
template <int HEAD_DIM, int BC, int STAGES = 2>
__global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
constexpr int BR = 16;
constexpr int KD = HEAD_DIM / 16;
constexpr int NC8 = BC / 8;
constexpr int KT2 = BC / 16;
constexpr int DN8 = HEAD_DIM / 8;
constexpr int LD = HEAD_DIM;
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
constexpr int VEC = 8;
constexpr int TOTAL = BC * HEAD_DIM;
const int lane = threadIdx.x;
const int gid = lane >> 2;
@@ -49,27 +42,19 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
const int G = p.q_head / p.kv_head;
const int q_head0 = kv_head * G;
__shared__ __align__(16) bf16 sK[BC * HEAD_DIM];
__shared__ __align__(16) bf16 sV[BC * HEAD_DIM];
__shared__ __align__(16) bf16 sQ[BR * HEAD_DIM];
for (int i = lane; i < BR * HEAD_DIM; i += 32) {
int r = i / HEAD_DIM, d = i % HEAD_DIM;
bf16 val = __float2bfloat16(0.0f);
if (r < G) {
int qh = q_head0 + r;
val = p.q[(batch * p.q_head + qh) * HEAD_DIM + d];
}
sQ[r * LD + swiz_col(d, r, SWIZ_MASK)] = val;
}
__syncwarp();
// 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];
int qrow_l = (lane & 7) + (lane & 8);
int qcol_l = (lane & 16) ? 8 : 0;
#pragma unroll
for (int kt = 0; kt < KD; kt++)
ldmatrix_x4(Qa[kt], &sQ[qrow_l * LD + swiz_col(kt * 16 + qcol_l, qrow_l, SWIZ_MASK)]);
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
@@ -77,60 +62,80 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
const int kv_base = (batch * p.kv_head + kv_head) * p.kv_len * HEAD_DIM;
const int mask_base = batch * p.kv_len;
// 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;
bool full_tile = (kv0 + BC <= p.kv_len);
if (full_tile) {
constexpr int VEC = 8;
int total = BC * HEAD_DIM;
#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;
cp_async_16(&sK[r * LD + swiz_col(d, r, SWIZ_MASK)],
&p.k[kv_base + kc * HEAD_DIM + d]);
cp_async_16(&sV[r * LD + swiz_col(d, r, SWIZ_MASK)],
&p.v[kv_base + kc * HEAD_DIM + d]);
}
cp_async_commit();
cp_async_wait_all();
} else {
for (int i = lane; i < BC * HEAD_DIM; i += 32) {
int r = i / HEAD_DIM, d = i % HEAD_DIM;
int kc = kv0 + r;
bf16 z = __float2bfloat16(0.0f);
sK[r * LD + swiz_col(d, r, SWIZ_MASK)] =
(kc < p.kv_len) ? p.k[kv_base + kc * HEAD_DIM + d] : z;
sV[r * LD + swiz_col(d, r, SWIZ_MASK)] =
(kc < p.kv_len) ? p.v[kv_base + kc * HEAD_DIM + d] : z;
}
}
__syncwarp();
float Sacc[NC8][4];
mma_compute_scores<KD, NC8>(Qa, sK, LD, SWIZ_MASK, lane, Sacc);
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
#pragma unroll
for (int n8 = 0; n8 < NC8; n8++)
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
int maxc = p.is_causal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
// 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,
mask_base, p.mask, has_mask,
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, sV, LD, SWIZ_MASK, lane, Oacc);
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 ----
@@ -169,4 +174,3 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
}
}
}
+152 -18
View File
@@ -1,43 +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,
bool is_causal,
int64_t causal_offset,
c10::optional<double> scale,
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.batch = (int)q.size(0);
p.q_head = (int)q.size(1);
p.kv_head = (int)k.size(1);
p.q_len = (int)q.size(2);
p.kv_len = (int)k.size(2);
p.head_dim = (int)q.size(3);
p.use_mask = mask.has_value() ? 1 : 0;
p.is_causal = is_causal ? 1 : 0;
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.scale = scale.has_value() ? (float)scale.value() : 1.0f / sqrtf((float)p.head_dim);
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();
if (p.use_mask) {
TORCH_CHECK(mask.value().dtype() == torch::kBool);
TORCH_CHECK(mask.value().dim() == 2);
TORCH_CHECK(mask.value().size(0) == p.batch);
TORCH_CHECK(mask.value().size(1) == p.kv_len);
p.mask = mask.value().data_ptr<bool>();
} else {
p.mask = nullptr;
}
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);
}
+74 -18
View File
@@ -108,20 +108,53 @@ __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 and pre-scaled by the caller).
// 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,
int LD,
int SWIZ_MASK,
int lane,
float Sacc[NC8][4])
{
#pragma unroll
@@ -138,15 +171,23 @@ __device__ inline void mma_compute_scores(
}
}
// Online softmax + Oacc rescale for one K/V tile. maxc0/maxc1 are the per-row
// KV column bounds prefill passes per-query-row causal limits while decode
// passes the same value for both rows (q_len==1). Sacc is consumed in place
// (replaced by P = exp(S - nm) for the subsequent P@V step).
// 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 mask_base,
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],
@@ -157,15 +198,19 @@ __device__ inline void mma_softmax_tile(
{
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_base + cc]);
bool b1 = (c1 >= maxc0) || (has_mask && !mask[mask_base + c1]);
bool b2 = (cc >= maxc1) || (has_mask && !mask[mask_base + cc]);
bool b3 = (c1 >= maxc1) || (has_mask && !mask[mask_base + c1]);
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];
@@ -175,22 +220,33 @@ __device__ inline void mma_softmax_tile(
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);
float corr0 = (nm0 == -FLT_MAX) ? 1.0f : __expf(m0 - nm0);
float corr1 = (nm1 == -FLT_MAX) ? 1.0f : __expf(m1 - nm1);
// 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 = (Sacc[n8][0] == -FLT_MAX) ? 0.0f : __expf(Sacc[n8][0] - nm0);
float p1 = (Sacc[n8][1] == -FLT_MAX) ? 0.0f : __expf(Sacc[n8][1] - nm0);
float p2 = (Sacc[n8][2] == -FLT_MAX) ? 0.0f : __expf(Sacc[n8][2] - nm1);
float p3 = (Sacc[n8][3] == -FLT_MAX) ? 0.0f : __expf(Sacc[n8][3] - nm1);
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;
+24 -94
View File
@@ -3,49 +3,29 @@
#include "attn_paged_decode_split_kv_mma.cuh"
#endif
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
static int paged_decode_num_splits(int base_blocks, int tiles_total) {
int sm_count = 0;
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
return std::max(1, std::min(n, std::min(tiles_total, 32)));
}
#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 = paged_decode_num_splits(p.batch * p.kv_head, chunks_total);
auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
auto o_part = torch::empty({p.batch, p.q_head, p.num_splits, p.head_dim}, fopt);
auto ml_part = torch::empty({p.batch, p.q_head, p.num_splits, 2}, fopt);
p.o_part = o_part.data_ptr<float>();
p.ml_part = ml_part.data_ptr<float>();
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);
paged_attn_decode_split_kv_kernel<<<
dim3(p.batch * p.kv_head, 1, p.num_splits),
dim3(32, group_size),
smem>>>(p);
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>
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
static void launch_paged_mma_decode(PagedAttentionParams<bf16>& p) {
int tiles_total = (p.kv_len + BC - 1) / BC;
p.num_splits = paged_decode_num_splits(p.batch * p.kv_head, tiles_total);
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
alloc_split_partials(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 = o_part.data_ptr<float>();
p.ml_part = ml_part.data_ptr<float>();
paged_attn_decode_split_kv_mma_kernel<HEAD_DIM, BC>
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(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
@@ -54,7 +34,7 @@ template <int HEAD_DIM>
static void dispatch_paged_decode(PagedAttentionParams<bf16>& p) {
#ifndef ASTRAI_NO_MMA
int G = p.q_head / p.kv_head;
if (!p.use_mask && G >= 1 && G <= 16 && p.page_size >= 32) {
if (G >= 1 && G <= 16 && p.page_size >= 32) {
launch_paged_mma_decode<HEAD_DIM, 32>(p);
return;
}
@@ -70,69 +50,19 @@ torch::Tensor attn_paged_decode(
int64_t page_size,
int64_t kv_len,
c10::optional<torch::Tensor> mask,
bool is_causal = false,
int64_t causal_offset = 0,
c10::optional<double> scale = c10::nullopt
int64_t causal_offset,
double scale,
int64_t layout
) {
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
PagedAttentionParams<bf16> p;
attn_pack_paged_params(q, page_table, k_cache, v_cache,
page_size, kv_len, mask, causal_offset, scale, layout, p);
int batch = q.size(0);
int q_head = q.size(1);
int head_dim = q.size(3);
int kv_head = k_cache.size(2);
int max_pages = page_table.size(1);
TORCH_CHECK(q.is_cuda() && page_table.is_cuda() && k_cache.is_cuda() && v_cache.is_cuda());
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16");
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16");
TORCH_CHECK(page_table.dtype() == torch::kLong, "page_table must be int64");
TORCH_CHECK(q.size(2) == 1, "Q seq_len must be 1 (decode)");
TORCH_CHECK(head_dim % 32 == 0, "head_dim must be multiple of 32");
TORCH_CHECK(k_cache.size(1) == page_size,
"k_cache dim 1 must equal page_size, got ",
k_cache.size(1), " vs ", page_size);
TORCH_CHECK(k_cache.size(0) >= 0, "k_cache must have at least 0 pages");
float scale_val = scale.has_value()
? static_cast<float>(scale.value())
: 1.0f / std::sqrt(static_cast<float>(head_dim));
auto O = torch::empty_like(q);
PagedAttentionParams<bf16, float> p;
p.batch = batch;
p.q_head = q_head;
p.kv_head = kv_head;
p.q_len = 1;
p.kv_len = static_cast<int>(kv_len);
p.head_dim = head_dim;
p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0;
p.is_causal = is_causal ? 1 : 0;
p.causal_offset = static_cast<int>(causal_offset);
p.num_splits = 1;
p.scale = scale_val;
p.page_size = static_cast<int>(page_size);
p.max_pages = max_pages;
p.page_table = page_table.data_ptr<int64_t>();
p.k_cache = reinterpret_cast<const bf16*>(k_cache.data_ptr());
p.v_cache = reinterpret_cast<const bf16*>(v_cache.data_ptr());
p.q = reinterpret_cast<const bf16*>(q.data_ptr());
p.mask = p.use_mask ? mask.value().data_ptr<bool>() : nullptr;
p.o = reinterpret_cast<bf16*>(O.data_ptr());
p.o_part = nullptr;
p.ml_part = nullptr;
switch (p.head_dim) {
case 32: dispatch_paged_decode<32>(p); break;
case 64: dispatch_paged_decode<64>(p); break;
case 128: dispatch_paged_decode<128>(p); break;
case 256: dispatch_paged_decode<256>(p); break;
default:
TORCH_CHECK(false, "paged_decode: unsupported head_dim ", p.head_dim,
" (supported: 32, 64, 128, 256)");
}
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;
}
@@ -145,8 +75,8 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
py::arg("page_size"),
py::arg("kv_len"),
py::arg("mask") = py::none(),
py::arg("is_causal") = false,
py::arg("causal_offset") = 0,
py::arg("scale") = py::none(),
py::arg("causal_offset") = -1,
py::arg("scale") = 0.0,
py::arg("layout") = 0,
"Paged GQA decode — split-KV with direct page-table access.");
}
+13 -6
View File
@@ -22,11 +22,13 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
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_head + q_head) * 1) * p.head_dim + lane * hd_per_thread;
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]);
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};
@@ -37,7 +39,7 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
int ch_begin = split * chunks_per_split;
int ch_end = min(chunks_total, ch_begin + chunks_per_split);
const int mask_base = batch * p.kv_len;
const int mask_base = batch * p.mask_b_stride;
for (int ci = ch_begin; ci < ch_end; ci++) {
int chunk_start = ci * PDC_CHUNK;
@@ -70,9 +72,10 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
partial += q_reg[i] * __bfloat162float(k_smem[s * p.head_dim + lane * hd_per_thread + i]);
partial = paged_warp_reduce_sum(partial) * p.scale;
if (p.use_mask && p.mask && !p.mask[mask_base + chunk_start + s])
int kv_idx = chunk_start + s;
if (p.use_mask && p.mask && !p.mask[mask_base + kv_idx])
partial = -FLT_MAX;
if (p.is_causal && (chunk_start + s) > p.causal_offset)
if (p.causal_offset >= 0 && kv_idx > p.causal_offset)
partial = -FLT_MAX;
float new_m = fmaxf(m, partial);
@@ -118,6 +121,9 @@ __global__ void paged_attn_decode_combine_kernel(PagedAttentionParams<bf16> p) {
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;
@@ -136,5 +142,6 @@ __global__ void paged_attn_decode_combine_kernel(PagedAttentionParams<bf16> p) {
}
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
p.o[(size_t)bh * p.head_dim + d] = __float2bfloat16(acc * inv);
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);
}
+69 -73
View File
@@ -11,56 +11,46 @@ using bf16 = __nv_bfloat16;
// 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>
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
__global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16> p) {
constexpr int BR = 16;
constexpr int KD = HEAD_DIM / 16;
constexpr int NC8 = BC / 8;
constexpr int KT2 = BC / 16;
constexpr int DN8 = HEAD_DIM / 8;
constexpr int LD = HEAD_DIM;
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
constexpr int VEC = 8;
constexpr int TOTAL = BC * HEAD_DIM;
const int lane = threadIdx.x;
const int gid = lane >> 2;
const int tid4 = lane & 3;
const int kv_head_idx = blockIdx.x;
const int 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_idx * G;
const int q_head0 = kv_head * G;
__shared__ __align__(16) bf16 sK[BC * HEAD_DIM];
__shared__ __align__(16) bf16 sV[BC * HEAD_DIM];
__shared__ __align__(16) bf16 sQ[BR * HEAD_DIM];
// ---- load Q into registers via ldmatrix ----
for (int i = lane; i < BR * HEAD_DIM; i += 32) {
int r = i / HEAD_DIM, d = i % HEAD_DIM;
bf16 val = __float2bfloat16(0.0f);
if (r < G) {
int qh = q_head0 + r;
val = p.q[(batch * p.q_head + qh) * HEAD_DIM + d];
}
sQ[r * LD + swiz_col(d, r, SWIZ_MASK)] = val;
}
__syncwarp();
__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];
int qrow_l = (lane & 7) + (lane & 8);
int qcol_l = (lane & 16) ? 8 : 0;
#pragma unroll
for (int kt = 0; kt < KD; kt++)
ldmatrix_x4(Qa[kt], &sQ[qrow_l * LD + swiz_col(kt * 16 + qcol_l, qrow_l, SWIZ_MASK)]);
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
#pragma unroll
for (int j = 0; j < DN8; j++)
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
const int mask_base = batch * p.kv_len;
const int tiles_total = (p.kv_len + BC - 1) / BC;
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
const int ti_begin = split * tiles_per_split;
@@ -70,70 +60,76 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
// Paged strides (constant for the block)
const int64_t page_stride = (int64_t)p.page_size * p.kv_head * HEAD_DIM;
const int64_t pos_stride = (int64_t)p.kv_head * HEAD_DIM;
const int64_t head_off = (int64_t)kv_head_idx * HEAD_DIM;
const int64_t head_off = (int64_t)kv_head * HEAD_DIM;
for (int ti = ti_begin; ti < ti_end; ti++) {
// ---- Load tile lambda: predicated cp.async, paged addressing ----
auto load_tile = [&](int ti, int buf) {
int kv0 = ti * BC;
// phys_page is constant for the whole tile (BC <= page_size).
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);
bool full_tile = page_valid && (kv0 + BC <= p.kv_len);
if (full_tile) {
constexpr int VEC = 8;
int total = BC * HEAD_DIM;
#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;
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;
cp_async_16(&sK[r * LD + swiz_col(d, r, SWIZ_MASK)],
&p.k_cache[gmem_base + d]);
cp_async_16(&sV[r * LD + swiz_col(d, r, SWIZ_MASK)],
&p.v_cache[gmem_base + d]);
}
cp_async_commit();
cp_async_wait_all();
} else {
for (int i = lane; i < BC * HEAD_DIM; i += 32) {
int r = i / HEAD_DIM, d = i % HEAD_DIM;
int kc = kv0 + r;
bf16 z = __float2bfloat16(0.0f);
if (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;
sK[r * LD + swiz_col(d, r, SWIZ_MASK)] = p.k_cache[gmem_base + d];
sV[r * LD + swiz_col(d, r, SWIZ_MASK)] = p.v_cache[gmem_base + d];
} else {
sK[r * LD + swiz_col(d, r, SWIZ_MASK)] = z;
sV[r * LD + swiz_col(d, r, SWIZ_MASK)] = z;
}
}
#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, sK, LD, SWIZ_MASK, lane, Sacc);
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
#pragma unroll
#pragma unroll
for (int n8 = 0; n8 < NC8; n8++)
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
int maxc = p.is_causal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
// 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,
mask_base, p.mask, has_mask,
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, sV, LD, SWIZ_MASK, lane, Oacc);
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 ----
@@ -141,7 +137,7 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
size_t bh = (size_t)batch * p.q_head + h;
return bh * p.num_splits + split;
};
#pragma unroll
#pragma unroll
for (int dn8 = 0; dn8 < DN8; dn8++) {
int d = dn8 * 8 + 2 * tid4;
int r0 = gid, r1 = gid + 8;
+17 -34
View File
@@ -13,17 +13,15 @@ static void dispatch_prefill(AttentionParams<bf16>& p) {
// 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. Both keep 3 blocks/SM (2 for D=256).
// over the 48KB static cap.
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
// Register-hint MIN_BLOCKS tuned per HEAD_DIM's (BC=32) smem+register
// footprint: the largest blocks/SM that avoids register spills.
constexpr int MIN_BLOCKS = (HEAD_DIM <= 32) ? 6 : (HEAD_DIM <= 64) ? 4
: (HEAD_DIM <= 128) ? 3 : 2;
dim3 grid((p.q_len + BR * WARPS - 1) / (BR * WARPS), p.q_head, p.batch);
dim3 block(WARPS * 32, 1, 1);
// Static shared memory — no dynamic smem or cudaFuncSetAttribute needed.
// sK[BC*LD] + sV[BC*LD] + sQ[BR*LD], all sized by template params.
attn_prefill_split_q_mma_kernel<HEAD_DIM, WARPS, BC, MIN_BLOCKS><<<grid, block>>>(p);
// 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);
@@ -37,34 +35,19 @@ torch::Tensor attn_prefill(
torch::Tensor k,
torch::Tensor v,
c10::optional<torch::Tensor> mask,
bool is_causal = false,
int64_t causal_offset = 0,
c10::optional<double> scale = c10::nullopt
int64_t causal_offset,
double scale,
int64_t layout
) {
AttentionParams<bf16> p;
attn_pack_params(q, k, v, mask, is_causal, causal_offset, scale, 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_like(q);
p.o = (bf16*)O.data_ptr();
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();
switch (p.head_dim) {
case 32:
dispatch_prefill<32>(p);
break;
case 64:
dispatch_prefill<64>(p);
break;
case 128:
dispatch_prefill<128>(p);
break;
case 256:
dispatch_prefill<256>(p);
break;
default:
TORCH_CHECK(false, "prefill: unsupported head_dim ", p.head_dim,
" (supported: 32,64,128,256)");
}
DISPATCH_HEAD_DIM(p.head_dim, dispatch_prefill, p);
return O;
}
@@ -74,8 +57,8 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
py::arg("k"),
py::arg("v"),
py::arg("mask") = py::none(),
py::arg("is_causal") = false,
py::arg("causal_offset") = 0,
py::arg("scale") = py::none(),
py::arg("causal_offset") = -1,
py::arg("scale") = 0.0,
py::arg("layout") = 0,
"GQA prefill (tensor-core mma on sm_80+, scalar fallback)");
}
+22 -10
View File
@@ -50,12 +50,14 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
__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_head + q_head) * p.q_len + q_row) * HEAD_DIM + gpos * DPT;
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.scale;
qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]) * p.scale;
}
float m = -FLT_MAX, l = 0.0f;
@@ -64,7 +66,9 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
for (int i = 0; i < DPT; i++)
acc[i] = 0.0f;
int kv_base = ((batch * p.kv_head + kv_head) * p.kv_len) * HEAD_DIM;
// 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;
@@ -79,15 +83,19 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
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 gidx = kv_base + (kv0 + i / HEAD_DIM) * HEAD_DIM + (i % HEAD_DIM);
sK[i] = p.k[gidx];
sV[i] = p.v[gidx];
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.is_causal && q_row < p.q_len) {
if (p.causal_offset >= 0 && q_row < p.q_len) {
int ep = q_row + p.causal_offset + 1;
if (kv0 >= ep)
lim = 0;
@@ -95,6 +103,7 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
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;
@@ -108,7 +117,8 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
}
float dot = group_reduce_sum<G>(part, gmask);
if (p.use_mask && p.mask && !p.mask[batch * p.kv_len + kv0 + s])
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);
@@ -131,10 +141,12 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
}
if (q_row < p.q_len) {
int o_off = ((batch * p.q_head + q_head) * p.q_len + q_row) * HEAD_DIM + gpos * DPT;
// 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] = __float2bfloat16(acc[i] * rl);
p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * rl);
}
}
+33 -46
View File
@@ -16,12 +16,7 @@ using bf16 = __nv_bfloat16;
// 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, MIN_BLOCKS> with BC a
// multiple of 16.
//
// Occupancy: __launch_bounds__ forces the compiler to fit MIN_BLOCKS blocks/SM,
// spilling to local memory as needed. MIN_BLOCKS is tuned per HEAD_DIM to the
// double-buffered smem footprint (2*BC*LD for each of K/V).
// 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
@@ -35,14 +30,14 @@ using bf16 = __nv_bfloat16;
// 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); pre-scale Q by attention scale during Q load; packed bf16x2 output stores;
// (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, int MIN_BLOCKS>
__global__ __launch_bounds__(WARPS * 32, MIN_BLOCKS)
void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
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)
@@ -62,7 +57,7 @@ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
const int kv_head = q_head / (p.q_head / p.kv_head);
const int qrow0 = (blockIdx.x * WARPS + warp) * BR;
// Static shared memory — sized by template parameters at compile time.
// ---- 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.
@@ -70,30 +65,16 @@ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
__shared__ __align__(16) bf16 sK[STAGES * BC * LD];
__shared__ __align__(16) bf16 sV[STAGES * BC * LD];
// Load the Q fragments straight from global into the mma A-operand layout
// (m16n8k16, row-major): no sQ staging area and no serialized per-warp
// prologue barriers. Each lane reads exactly the 8 Q elements ldmatrix
// would have produced, pre-scaled by the attention scale. Kept resident in
// registers across the tile loop.
// frag[0]/[2]: row = qrow0 + gid ; frag[1]/[3]: row = qrow0 + gid + 8
// frag[0]/[1]: cols kt*16 + tid4*2 + {0,1} ; frag[2]/[3]: + 8
const int q_base = ((batch * p.q_head + q_head) * p.q_len) * HEAD_DIM;
// 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];
#pragma unroll
for (int kt = 0; kt < KD; kt++) {
int c = kt * 16 + tid4 * 2;
const unsigned* pau = reinterpret_cast<const unsigned*>(
&p.q[q_base + qra * HEAD_DIM + c]);
const unsigned* pbu = reinterpret_cast<const unsigned*>(
&p.q[q_base + qrb * HEAD_DIM + c]);
Qa[kt][0] = va ? pau[0] : 0u;
Qa[kt][1] = vb ? pbu[0] : 0u;
Qa[kt][2] = va ? pau[4] : 0u;
Qa[kt][3] = vb ? pbu[4] : 0u;
}
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
@@ -101,18 +82,18 @@ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
const int kv_base = ((batch * p.kv_head + kv_head) * p.kv_len) * HEAD_DIM;
// 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 is_causal == 0)
const int use_skip = p.is_causal;
// 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;
const int mb = batch * p.kv_len;
// Last active tile: block-level causal bound (all warps in the block share
// the K/V load, so the prefetch range is the block max, not per-warp).
@@ -125,6 +106,7 @@ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
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) {
@@ -137,13 +119,14 @@ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
int kc = kv0 + r;
bool valid = kc < p.kv_len;
int off = r * LD + swiz_col(d, r, SWIZ_MASK);
cp_async_16_pred(&dK[off], &p.k[kv_base + kc * HEAD_DIM + d], valid);
cp_async_16_pred(&dV[off], &p.v[kv_base + kc * HEAD_DIM + d], valid);
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: kick off the first tile's load.
// ---- Prologue: issue first tile load ----
load_tile(0, 0);
for (int ti = 0; ti <= t_end; ti++) {
@@ -175,12 +158,15 @@ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
int maxc0 = p.is_causal ? min(p.kv_len, qr0 + p.causal_offset + 1)
: p.kv_len;
int maxc1 = p.is_causal ? min(p.kv_len, qr1 + p.causal_offset + 1)
: p.kv_len;
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,
mb, p.mask, has_mask,
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);
@@ -191,19 +177,20 @@ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
// halves store count and removes the uncoalesced scalar-store penalty)
float rl0 = (l0 > 1e-20f) ? (1.0f / l0) : 0.0f;
float rl1 = (l1 > 1e-20f) ? (1.0f / l1) : 0.0f;
const int o_base = ((batch * p.q_head + q_head) * p.q_len) * HEAD_DIM;
// 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 * HEAD_DIM + d]) = v;
*reinterpret_cast<__nv_bfloat162*>(&p.o[o_base + qr0 * p.q_stride_l + d * p.q_stride_d]) = v;
}
if (qr1 < p.q_len) {
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1,
Oacc[dn8][3] * rl1);
*reinterpret_cast<__nv_bfloat162*>(&p.o[o_base + qr1 * HEAD_DIM + d]) = v;
*reinterpret_cast<__nv_bfloat162*>(&p.o[o_base + qr1 * p.q_stride_l + d * p.q_stride_d]) = v;
}
}
}
+37 -61
View File
@@ -27,14 +27,14 @@ static bool decode_use_mma(const AttentionParams<bf16>& p) {
return !p.use_mask && G > 1 && G <= 16;
}
template <int HEAD_DIM, int BC>
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>
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);
}
@@ -61,90 +61,65 @@ static void dispatch_decode_t(AttentionParams<bf16>& p, DecodeScratch& sc) {
}
static void dispatch_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
switch (p.head_dim) {
case 32: dispatch_decode_t<32>(p, sc); break;
case 64: dispatch_decode_t<64>(p, sc); break;
case 128: dispatch_decode_t<128>(p, sc); break;
case 256: dispatch_decode_t<256>(p, sc); break;
default: printf("bench: unsupported D=%d\n", p.head_dim);
}
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.
// Decode (q_len==1) is memory-bound: the two matmuls are GEMV-shaped, so we
// report both effective K/V read bandwidth and the (small) attention FLOP/s.
// FLOP/s = 2 matmuls (q@K^T, P@V), each 2*B*Hq*kv*D flops.
// Bytes = K + V read = 2 * B*Hk*kv*D * sizeof(bf16).
static void bench() {
const int cfgs[][5] = {
{1, 32, 4, 512, 128}, // B,Hq,Hk,seq,D
{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},
};
int n = sizeof(cfgs)/sizeof(cfgs[0]);
const int WARMUP = 10, ITERS = 100;
printf("\n===== DECODE BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
printf("%-46s | %10s | %10s | %10s\n",
"config", "latency", "bandwidth", "throughput");
printf("---------------------------------------------------------------"
"----------------------------\n");
print_bench_header();
for (int ci = 0; ci < n; 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, nKV=(size_t)B*Hk*sl*D;
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,*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);
bf16 *dQ, *dK, *dV, *dO;
cudaMalloc(&dQ, nQ*2); cudaMalloc(&dK, nKV*2);
cudaMalloc(&dV, nKV*2); cudaMalloc(&dO, nQ*2);
size_t big = nQ > nKV ? nQ : nKV; bf16* tmp = new bf16[big];
for (size_t i = 0; i < nQ; i++) tmp[i] = f2bf(randf());
cudaMemcpy(dQ, tmp, nQ*2, cudaMemcpyHostToDevice);
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
cudaMemcpy(dK, tmp, nKV*2, cudaMemcpyHostToDevice);
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
cudaMemcpy(dV, tmp, nKV*2, cudaMemcpyHostToDevice);
delete[] tmp;
AttentionParams<bf16> p;
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D;
p.use_mask=0; p.is_causal=0; p.causal_offset=0;
p.scale=1.0f/sqrtf((float)D);
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
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));
for (int i=0;i<WARMUP;i++) dispatch_decode(p, sc);
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_decode(p, sc);
cudaEventRecord(e); cudaEventSynchronize(e);
float ms=0; cudaEventElapsedTime(&ms,s,e); ms/=ITERS;
double flops = 4.0*B*Hq*(double)sl*D;
double tflops = flops/(ms*1e-3)/1e12;
// HBM traffic: K + V read (B*Hk*sl*D each), bf16; Q/O negligible.
double bytes = 2.0 * (2.0*nKV);
double gbps = bytes/(ms*1e-3)/1e9;
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);
printf("%-46s | %7.4f ms | %7.1f GB/s | %6.2f TFLOP/s\n",
cfg, ms, gbps, tflops);
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);
delete[]tmp; cudaEventDestroy(s); cudaEventDestroy(e);
cudaFree(dQ); cudaFree(dK); cudaFree(dV); cudaFree(dO);
cudaFree(sc.o_part); cudaFree(sc.ml_part);
}
}
@@ -186,8 +161,9 @@ int main() {
AttentionParams<bf16> p;
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D;
p.use_mask=0; p.is_causal=0; p.causal_offset=0;
p.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.
@@ -206,7 +182,7 @@ int main() {
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
float* ref=new float[nQ];
cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, 0, 0);
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++){
+19 -33
View File
@@ -42,9 +42,10 @@ static void launch_paged_decode(PagedAttentionParams<bf16, float>& p) {
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>
paged_attn_decode_split_kv_mma_kernel<HEAD_DIM, 32, STAGES>
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
} else
#endif
@@ -137,13 +138,14 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed)
}
float* h_o_ref = (float*)calloc(B * Hq * HEAD_DIM, sizeof(float));
cpu_attention_ref(h_q_f, h_k_f, h_v_f, nullptr, h_o_ref, B, Hq, Hkv, 1, kv_len, HEAD_DIM, 0, 0);
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.is_causal = 0; p.causal_offset = 0;
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;
@@ -221,17 +223,16 @@ static const TestCase TESTS[] = {
};
static int dispatch_test(const TestCase& tc) {
switch (tc.head_dim) {
case 32: return run_test<32>(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.seed);
case 64: return run_test<64>(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.seed);
case 128: return run_test<128>(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.seed);
case 256: return run_test<256>(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.seed);
default: return 1;
}
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.
// Reports per-call latency and effective K/V read bandwidth.
// 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) {
@@ -272,7 +273,8 @@ static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) {
PagedAttentionParams<bf16, float> pa;
pa.batch = B; pa.q_head = Hq; pa.kv_head = Hkv; pa.q_len = 1;
pa.kv_len = kv_len; pa.head_dim = HEAD_DIM;
pa.use_mask = 0; pa.is_causal = 0; pa.causal_offset = 0;
pa.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;
@@ -281,43 +283,27 @@ static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) {
pa.o_part = d_op; pa.ml_part = d_ml;
const int WARMUP = 10, ITERS = 100;
for (int i = 0; i < WARMUP; i++) launch_paged_decode<HEAD_DIM>(pa);
cudaDeviceSynchronize();
CUDA_CHECK(cudaGetLastError());
cudaEvent_t s, e;
cudaEventCreate(&s); cudaEventCreate(&e);
cudaEventRecord(s);
for (int i = 0; i < ITERS; i++) launch_paged_decode<HEAD_DIM>(pa);
cudaEventRecord(e); cudaEventSynchronize(e);
float ms = 0; cudaEventElapsedTime(&ms, s, e); ms /= ITERS;
auto launch = [&]() { launch_paged_decode<HEAD_DIM>(pa); };
double flops = 4.0 * B * Hq * (double)kv_len * HEAD_DIM;
double tflops = flops / (ms * 1e-3) / 1e12;
size_t nKV = (size_t)B * Hkv * kv_len * HEAD_DIM;
double bytes = 2.0 * (2.0 * nKV);
double gbps = bytes / (ms * 1e-3) / 1e9;
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);
printf("%-46s | %7.4f ms | %7.1f GB/s | %6.2f TFLOP/s\n",
cfg, ms, gbps, tflops);
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);
cudaEventDestroy(s); cudaEventDestroy(e);
}
static void bench() {
printf("\n===== PAGED DECODE BENCH =====\n");
printf("%-46s | %10s | %10s | %10s\n",
"config", "latency", "bandwidth", "throughput");
printf("---------------------------------------------------------------"
"----------------------------\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);
+6 -6
View File
@@ -18,11 +18,9 @@ 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;
constexpr int MIN_BLOCKS = (HEAD_DIM <= 32) ? 6 : (HEAD_DIM <= 64) ? 4
: (HEAD_DIM <= 128) ? 3 : 2;
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, MIN_BLOCKS><<<grid, block>>>(p);
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);
@@ -77,7 +75,8 @@ static void bench() {
AttentionParams<bf16> p;
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
p.use_mask=0; p.is_causal=causal; p.causal_offset=0;
p.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;
@@ -145,7 +144,8 @@ int main() {
AttentionParams<bf16> p;
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
p.use_mask=0; p.is_causal=causal; p.causal_offset=0;
p.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;
@@ -160,7 +160,7 @@ int main() {
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);
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++) {
+93 -2
View File
@@ -37,6 +37,96 @@ inline int compute_num_splits(int base_blocks, int tiles_total) {
} \
} 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]
@@ -44,10 +134,11 @@ inline int compute_num_splits(int base_blocks, int tiles_total) {
// 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 is_causal, int causal_offset
int causal_offset
) {
float scale = 1.0f / sqrtf((float)D);
int n_rep = Hq / Hk;
@@ -58,7 +149,7 @@ static void cpu_attention_ref(
float mv = -INFINITY, sv = 0.0f;
float accum[256] = {0.0f};
int lim = kv_len;
if (is_causal) {
if (causal_offset >= 0) {
int c = qi + causal_offset + 1;
lim = (c < kv_len) ? c : kv_len;
}
+15
View File
@@ -42,6 +42,19 @@ def parse_args():
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,
@@ -79,6 +92,8 @@ def chat():
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
+16 -2
View File
@@ -117,7 +117,7 @@ def print_component_summary(results: dict[str, dict], title: str):
r["er_99_norm"]
for vs in matrix_groups.values()
for r in vs
if "_norm" not in r or not r.get("is_1d")
if not r.get("is_1d")
]
if all_er:
m = sum(all_er) / len(all_er)
@@ -232,8 +232,16 @@ def main():
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"
@@ -294,13 +302,19 @@ def main():
)
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, "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__":
+38 -19
View File
@@ -8,7 +8,6 @@ Config is a single dataclass; side effects are isolated at pipeline boundaries.
"""
import argparse
import itertools
import json
import os
import re
@@ -233,7 +232,7 @@ def execute_one(args: tuple) -> bool:
return False
def test_one(item: dict, cfg: EvalConfig) -> Tuple[str, int, int]:
def test_one(item: dict, cfg: EvalConfig, pool=None) -> Tuple[str, int, int]:
from concurrent.futures import ProcessPoolExecutor
task_id = item["task_id"]
@@ -247,11 +246,16 @@ def test_one(item: dict, cfg: EvalConfig) -> Tuple[str, int, int]:
for c in completions
]
n = len(codes)
passed = 0
with ProcessPoolExecutor(max_workers=cfg.test_workers) as pool:
for ok in pool.map(execute_one, codes):
if ok:
passed += 1
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
@@ -259,8 +263,14 @@ def test_all(
items: Sequence[dict],
cfg: EvalConfig,
) -> Iterator[Tuple[str, int, int]]:
for item in tqdm.tqdm(items, desc="Testing", unit="problem"):
yield test_one(item, cfg)
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:
@@ -273,26 +283,32 @@ def score_results(
results: Iterator[Tuple[str, int, int]],
k_values: Tuple[int, ...],
) -> Dict:
# filter to k <= n (peek first result to get n)
first = next(results)
results = itertools.chain([first], results)
n = first[1]
k_values = tuple(k for k in k_values if k <= n)
"""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:
pk = round(pass_at_k(n, passed, k), 4)
entry[f"pass@{k}"] = pk
scores[k].append(pk)
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]
summary[f"pass@{k}"] = round(float(np.mean(vals)), 4)
if vals:
summary[f"pass@{k}"] = round(float(np.mean(vals)), 4)
else:
summary[f"pass@{k}"] = None
output["_summary"] = summary
return output
@@ -375,7 +391,10 @@ 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
+125 -87
View File
@@ -16,7 +16,9 @@ v2 changelog:
"""
import argparse
import glob
import json
import os
import statistics
import torch
@@ -65,6 +67,29 @@ def _resolve_sentinel_ids(tokenizer, sentinel_text):
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
@@ -204,69 +229,9 @@ def _trim(context_ids, resp_ids, max_len):
return context_ids[overflow:], resp_ids
def score_plain(
def process_file(
model,
tokenizer,
instruction,
response,
device,
max_len=2048,
sentinel_ids=None,
per_token=False,
):
"""Compute IFD for a single instruction-response pair (plain format)."""
ctx_ids = tokenizer.encode(instruction, add_special_tokens=False)
resp_ids = tokenizer.encode(response, add_special_tokens=False)
ctx_ids, resp_ids = _trim(ctx_ids, resp_ids, max_len)
if not ctx_ids or not resp_ids:
return {
"L_cond": None,
"L_uncond": None,
"ifd": None,
"skip_reason": "empty ctx or resp",
}
return _score_batch(
[(ctx_ids, resp_ids)],
model,
device,
max_len,
sentinel_ids=sentinel_ids,
per_token=per_token,
)[0]
def score_messages(
model, tokenizer, messages, device, max_len=2048, sentinel_ids=None, per_token=False
):
"""Compute IFD for each assistant turn in a messages array."""
turns = []
for i, msg in enumerate(messages):
if msg.get("role") != "assistant":
continue
ctx_text = "\n\n".join(m["content"] for m in messages[:i])
ctx_ids = tokenizer.encode(ctx_text)
resp_ids = tokenizer.encode(msg["content"], add_special_tokens=False)
ctx_ids, resp_ids = _trim(ctx_ids, resp_ids, max_len)
if ctx_ids and resp_ids:
turns.append((ctx_ids, resp_ids))
if not turns:
return None
raw_scores = _score_batch(
turns, model, device, max_len, sentinel_ids=sentinel_ids, per_token=per_token
)
valid = [s for s in raw_scores if s is not None and s.get("ifd") is not None]
if not valid:
return {"ifd": None, "ifd_turns": raw_scores}
avg = sum(s["ifd"] for s in valid) / len(valid)
return {
"ifd": avg,
"ifd_detail": valid[0] if len(valid) == 1 else None,
"ifd_turns": raw_scores,
}
def process_file(
param_path,
input_file,
output_file,
instr_key,
@@ -275,28 +240,31 @@ def process_file(
data_format="plain",
batch_size=1,
device=None,
sentinel_text="\n",
sentinel_ids=None,
per_token=False,
max_samples=None,
):
"""Score a single file, write per-sample JSONL, return summary stats."""
if device is None:
device = "cuda" if torch.cuda.is_available() else "cpu"
dtype = torch.bfloat16 if "cuda" in device else torch.float32
model = AutoModel.from_pretrained(param_path)
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")
sentinel_ids = _resolve_sentinel_ids(tokenizer, sentinel_text)
data = _load_items(input_file)
with open(input_file, 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 = []
all_ifds = []
buffer = []
for item in tqdm.tqdm(data, desc="Computing IFD", unit="sample"):
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", [])):
@@ -356,8 +324,22 @@ def process_file(
f.write(json.dumps(item, ensure_ascii=False) + "\n")
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:
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" [{label}]")
print(f"{'=' * 50}")
print(f" Samples: {len(data)}")
print(f" Valid IFD: {len(valid_ifd)}")
print(f" Skipped: {len(data) - len(valid_ifd)}")
@@ -368,7 +350,8 @@ def process_file(
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}")
print(f" Results saved to {output_file}")
return stats
def _flush_buffer(
@@ -422,8 +405,18 @@ 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(
"--input_path",
type=str,
required=True,
help="Input file, glob pattern, or directory.",
)
parser.add_argument(
"--output_dir",
type=str,
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",
@@ -442,6 +435,12 @@ def main():
"--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,
@@ -453,21 +452,60 @@ def main():
action="store_true",
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,
data_format=args.format,
batch_size=args.batch_size,
device=args.device,
sentinel_text=args.sentinel_text,
per_token=args.per_token,
)
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__":
+1 -1
View File
@@ -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
"""
+7 -14
View File
@@ -139,17 +139,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"
@@ -218,9 +213,7 @@ def evaluate_subject(
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 []
)
raw_prompt = build_prompt(item["question"], item, subject)
context = apply_chat(tokenizer, raw_prompt, n_shot, dev_data or [])
context_ids = tokenizer.encode(context)
scores = {
+422 -69
View File
@@ -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="Perplexity evaluation on JSONL text.")
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,
)
+10 -8
View File
@@ -90,6 +90,7 @@ class MuonMix(optim.Optimizer):
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:
@@ -115,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."
@@ -252,12 +260,6 @@ def parse_args() -> argparse.Namespace:
default="checkpoint/logs",
help="Directory for 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."
)
@@ -390,6 +392,7 @@ def train(
train_type: str,
param_path: str,
data_root_path: str,
resume: bool,
n_epoch: int,
batch_per_device: int,
start_epoch: int,
@@ -444,7 +447,6 @@ def train(
"clip_eps": kwargs.pop("grpo_clip_eps"),
"kl_coef": kwargs.pop("grpo_kl_coef"),
"group_size": kwargs.pop("group_size"),
"sync_interval": kwargs.pop("grpo_sync_interval"),
}
executor_kwargs = {
@@ -537,7 +539,7 @@ def train(
)
trainer = Trainer(train_config)
trainer.train(resume_dir=param_path)
trainer.train(param_path=param_path, resume=resume)
if __name__ == "__main__":
+399 -23
View File
@@ -327,43 +327,82 @@ def test_normalize_mixed_empty_key():
def test_grpo_dataset_dtype(base_test_env):
"""GRPO dataset returns correct dtypes for per-record structured data."""
from astrai.dataset.dataset import GRPODataset
test_dir = base_test_env["test_dir"]
dummy_data = {
"prompts": [torch.randint(0, 100, (100,), dtype=torch.int32)],
"responses": [torch.randint(0, 100, (100,), dtype=torch.int32)],
"masks": [torch.ones(100, dtype=torch.int32)],
"rewards": [torch.ones(100, dtype=torch.float32)],
}
dataset = _make_seq_dataset(
test_dir, "grpo_dtype", train_type="grpo", data=dummy_data, window_size=32
)
G = 4
dataset = GRPODataset()
dataset.storage = type(
"FakeStore",
(),
{
"keys": ["prompts", "responses", "masks", "rewards"],
"_data": {
"prompts": [torch.randint(0, 100, (10,), dtype=torch.int32)],
"responses": [
[torch.randint(0, 100, (5,), dtype=torch.int32) for _ in range(G)]
],
"masks": [[torch.ones(5, dtype=torch.int32) for _ in range(G)]],
"rewards": [torch.rand(G, dtype=torch.float32)],
},
},
)()
dataset._build_records()
item = dataset[0]
assert item["prompts"].dtype == torch.long
assert item["responses"].dtype == torch.long
assert item["masks"].dtype == torch.bool
assert all(r.dtype == torch.long for r in item["responses"])
assert all(m.dtype == torch.bool for m in item["masks"])
assert item["rewards"].dtype == torch.float32
def test_grpo_dataset_load(base_test_env):
"""GRPO dataset loads record-structured data with per-response boundaries."""
from astrai.dataset.dataset import GRPODataset
test_dir = base_test_env["test_dir"]
dummy_data = {
"prompts": [_rand_seq(200)],
"responses": [_rand_seq(200)],
"masks": [torch.ones(200, dtype=torch.int64)],
"rewards": [torch.rand(200, dtype=torch.float32)],
}
dataset = _make_seq_dataset(
test_dir, "grpo_test", train_type="grpo", data=dummy_data
)
assert len(dataset) > 0
G = 3
prompt_len = 8
resp_lens = [5, 7, 4]
dataset = GRPODataset()
dataset.storage = type(
"FakeStore",
(),
{
"keys": ["prompts", "responses", "masks", "rewards"],
"_data": {
"prompts": [torch.randint(0, 100, (prompt_len,))],
"responses": [[torch.randint(0, 100, (rl,)) for rl in resp_lens]],
"masks": [[torch.ones(rl, dtype=torch.int64) for rl in resp_lens]],
"rewards": [torch.tensor([0.9, 0.3, 0.7], dtype=torch.float32)],
},
},
)()
dataset._build_records()
assert len(dataset) == 1
item = dataset[0]
assert "prompts" in item
assert "responses" in item
assert "masks" in item
assert "rewards" in item
assert item["prompts"].shape[0] == 64
assert item["responses"].shape[0] == 64
# Prompts is 1-D
assert item["prompts"].shape == (prompt_len,)
# Responses is a list of G tensors with correct lengths
assert len(item["responses"]) == G
for i, r in enumerate(item["responses"]):
assert r.shape == (resp_lens[i],)
# Masks align with responses
assert len(item["masks"]) == G
for i, m in enumerate(item["masks"]):
assert m.shape == (resp_lens[i],)
# Rewards has G elements
assert item["rewards"].shape == (G,)
def test_detect_format_bin_dir(base_test_env):
@@ -411,6 +450,32 @@ def test_dataset_load_explicit_storage_type(base_test_env):
assert dataset.count == 200
def _write_json_dataset(test_dir, tokenizer_path, records, config_overrides=None):
"""Write JSON (not JSONL) dataset — array of objects."""
data_dir = os.path.join(test_dir, "json_data")
os.makedirs(data_dir, exist_ok=True)
with open(os.path.join(data_dir, "data.json"), "w", encoding="utf-8") as f:
json.dump(records, f, ensure_ascii=False)
config = {
"tokenizer_path": tokenizer_path,
"version": 1,
"input": {"sections": [{"field": "text", "action": "train"}]},
"preprocessing": {"max_seq_len": 128, "min_chars": 0},
"output": {"position_ids_mode": "continuous"},
}
if config_overrides:
config.update(config_overrides)
with open(
os.path.join(data_dir, "dataset_config.json"), "w", encoding="utf-8"
) as f:
json.dump(config, f, ensure_ascii=False, indent=2)
return data_dir
def test_detect_format_jsonl_dir(base_test_env):
test_dir = base_test_env["test_dir"]
tokenizer_path = _save_test_tokenizer(test_dir, base_test_env["tokenizer"])
@@ -422,6 +487,89 @@ def test_detect_format_jsonl_dir(base_test_env):
assert detect_format(data_dir) == "jsonl"
def test_detect_format_json_dir(base_test_env):
"""detect_format returns 'jsonl' for directory with .json files."""
test_dir = base_test_env["test_dir"]
tokenizer_path = _save_test_tokenizer(test_dir, base_test_env["tokenizer"])
data_dir = _write_json_dataset(
test_dir,
tokenizer_path,
[{"text": "hello world"}, {"text": "foo bar baz qux"}],
)
assert detect_format(data_dir) == "jsonl"
def test_json_store_seq(base_test_env):
"""JsonlStore loads .json array correctly."""
test_dir = base_test_env["test_dir"]
tokenizer_path = _save_test_tokenizer(test_dir, base_test_env["tokenizer"])
data_dir = _write_json_dataset(
test_dir,
tokenizer_path,
[{"text": "hello world"}, {"text": "foo bar baz qux"}],
)
store = StoreFactory.create("jsonl")
store.load(data_dir)
assert len(store) > 0
assert "sequence" in store.keys
dataset = DatasetFactory.load("seq", data_dir, window_size=8)
assert len(dataset) > 0
item = dataset[0]
assert "input_ids" in item
assert "target_ids" in item
def test_json_store_no_tokenizer_path(base_test_env):
"""JsonlStore uses dataset dir as tokenizer_path when omitted."""
test_dir = base_test_env["test_dir"]
tokenizer = base_test_env["tokenizer"]
tokenizer.set_chat_template(
"{% for message in messages %}{{ message['role'] }}:{{ message['content'] }}\n{% endfor %}"
)
data_dir = os.path.join(test_dir, "self_contained")
os.makedirs(data_dir, exist_ok=True)
# Save tokenizer files directly in the dataset directory
tokenizer.save_pretrained(data_dir)
# Write .json data
records = [
{
"messages": [
{"role": "user", "content": "hi"},
{"role": "assistant", "content": "hello"},
]
}
]
with open(os.path.join(data_dir, "data.json"), "w", encoding="utf-8") as f:
json.dump(records, f, ensure_ascii=False)
# dataset_config.json WITHOUT tokenizer_path
config = {
"version": 1,
"input": {
"sections": [{"field": "messages", "action": "$role", "template": True}]
},
"mask": {"user": "mask", "assistant": "train"},
"mask_default": "mask",
"preprocessing": {"max_seq_len": 128, "min_chars": 0},
"output": {"position_ids_mode": "continuous"},
}
with open(
os.path.join(data_dir, "dataset_config.json"), "w", encoding="utf-8"
) as f:
json.dump(config, f, ensure_ascii=False, indent=2)
store = StoreFactory.create("jsonl")
store.load(data_dir)
assert len(store) > 0
assert "sequence" in store.keys
assert "loss_mask" in store.keys
def test_jsonl_store_seq(base_test_env):
test_dir = base_test_env["test_dir"]
tokenizer_path = _save_test_tokenizer(test_dir, base_test_env["tokenizer"])
@@ -512,3 +660,231 @@ def test_jsonl_store_pipeline_config_roundtrip(base_test_env):
config = PipelineConfig.from_dict(raw)
assert config.output.position_ids_mode == "doc_reset"
assert config.preprocessing.max_seq_len == 64
# ---------------------------------------------------------------------------
# GRPO end-to-end: builder → JsonlStore → GRPODataset → collate_fn
# ---------------------------------------------------------------------------
def _write_grpo_jsonl(test_dir, tokenizer_path, records):
"""Write a GRPO JSONL dataset directory with config."""
data_dir = os.path.join(test_dir, "grpo_jsonl")
os.makedirs(data_dir, exist_ok=True)
with open(os.path.join(data_dir, "data.jsonl"), "w", encoding="utf-8") as f:
for rec in records:
f.write(json.dumps(rec, ensure_ascii=False) + "\n")
config = {
"tokenizer_path": tokenizer_path,
"version": 1,
"input": {
"sources": {
"prompts": {
"sections": [
{
"field": "prompt",
"action": "mask",
"add_special_tokens": True,
}
]
},
"responses": {
"sections": [{"field": "responses", "action": "train"}],
"list_field": True,
"mask_key": "masks",
},
"rewards": {
"sections": [{"field": "rewards", "action": "value"}],
},
}
},
"mask": {"user": "mask", "assistant": "train"},
"mask_default": "mask",
"preprocessing": {"max_seq_len": 128},
"output": {"position_ids_mode": "none"},
}
with open(
os.path.join(data_dir, "dataset_config.json"), "w", encoding="utf-8"
) as f:
json.dump(config, f, ensure_ascii=False, indent=2)
return data_dir
def test_grpo_builder_preserves_response_boundaries(base_test_env):
"""MultiOutputMaskBuilder with list_field returns List[List[int]] for responses."""
from astrai.preprocessing.builder import SectionedMaskBuilder
from tests.data.conftest import make_grpo_no_template_config
tokenizer = base_test_env["tokenizer"]
tokenizer_path = _save_test_tokenizer(base_test_env["test_dir"], tokenizer)
builder = SectionedMaskBuilder()
config = make_grpo_no_template_config()
config.preprocessing.max_seq_len = 128
item = {
"prompt": "What is 2+2?",
"responses": ["4", "four", "2+2=4"],
"rewards": [0.9, 0.1, 0.5],
}
result = builder.build(item, config, tokenizer)
assert result is not None
# prompts should be flat list of ints
assert isinstance(result["prompts"], list)
assert isinstance(result["prompts"][0], int)
# responses should be list of lists (one per response)
assert isinstance(result["responses"], list)
assert isinstance(result["responses"][0], list)
assert isinstance(result["responses"][0][0], int)
assert len(result["responses"]) == 3
# masks should match responses structure
assert isinstance(result["masks"], list)
assert len(result["masks"]) == 3
for i in range(3):
assert len(result["masks"][i]) == len(result["responses"][i])
# rewards should be flat list of floats
assert isinstance(result["rewards"], list)
assert all(isinstance(r, float) for r in result["rewards"])
assert len(result["rewards"]) == 3
def test_grpo_end_to_end_jsonl(base_test_env):
"""Full GRPO pipeline: JSONL → JsonlStore → GRPODataset → collate_fn."""
from astrai.dataset.dataset import grpo_collate_fn
test_dir = base_test_env["test_dir"]
tokenizer = base_test_env["tokenizer"]
tokenizer_path = _save_test_tokenizer(test_dir, tokenizer)
records = [
{
"prompt": "What is 2+2?",
"responses": ["4", "four", "The answer is 4"],
"rewards": [0.9, 0.1, 0.5],
},
{
"prompt": "Write a haiku",
"responses": ["Leaves fall", "Cherry blossoms bloom in spring"],
"rewards": [0.3, 0.8],
},
]
data_dir = _write_grpo_jsonl(test_dir, tokenizer_path, records)
dataset = DatasetFactory.load("grpo", data_dir, window_size=0)
assert len(dataset) == 2
# Item 0: 3 responses
item0 = dataset[0]
assert item0["prompts"].ndim == 1
assert len(item0["responses"]) == 3
assert len(item0["masks"]) == 3
assert item0["rewards"].shape == (3,)
for r, m in zip(item0["responses"], item0["masks"]):
assert r.shape == m.shape
# Item 1: 2 responses (different group size)
item1 = dataset[1]
assert len(item1["responses"]) == 2
assert item1["rewards"].shape == (2,)
# Collate: batch records with same G (item0 has G=3)
batch = grpo_collate_fn([item0, item0])
assert batch["prompts"].shape[0] == 2
assert batch["responses"].ndim == 3
assert batch["responses"].shape[0] == 2
assert batch["responses"].shape[1] == 3 # G=3
assert batch["masks"].shape == batch["responses"].shape
assert batch["rewards"].shape == (2, 3)
def test_grpo_collate_variable_lengths():
"""collate_fn pads variable-length responses to [B, G, R_max]."""
from astrai.dataset.dataset import grpo_collate_fn
batch = [
{
"prompts": torch.tensor([1, 2, 3]),
"responses": [torch.tensor([4, 5]), torch.tensor([6, 7, 8, 9])],
"masks": [torch.tensor([1, 1]), torch.tensor([1, 1, 1, 1])],
"rewards": torch.tensor([0.9, 0.1]),
},
{
"prompts": torch.tensor([10, 11]),
"responses": [torch.tensor([12]), torch.tensor([13, 14, 15])],
"masks": [torch.tensor([1]), torch.tensor([1, 1, 1])],
"rewards": torch.tensor([0.5, 0.5]),
},
]
result = grpo_collate_fn(batch)
assert result["prompts"].shape == (2, 3) # B=2, P_max=3
assert result["responses"].shape == (2, 2, 4) # B=2, G=2, R_max=4
assert result["masks"].shape == (2, 2, 4)
assert result["rewards"].shape == (2, 2)
# Check padding: item 1 prompt is length 2, padded to 3
assert result["prompts"][1, 2] == 0
# Check response content: item 0, response 0 is [4,5] padded to 4
assert result["responses"][0, 0, 0] == 4
assert result["responses"][0, 0, 1] == 5
assert result["responses"][0, 0, 2] == 0 # padded
assert result["masks"][0, 0, 2] == False # padded
# Check response content: item 0, response 1 is [6,7,8,9] no padding
assert result["responses"][0, 1, 3] == 9
assert result["masks"][0, 1, 3] == True
def test_grpo_multiple_records(base_test_env):
"""GRPODataset loads multiple records with correct structure."""
from astrai.dataset.dataset import GRPODataset
G = 4
n_records = 5
dummy_responses = [
[torch.randint(0, 100, (np.random.randint(3, 8),)) for _ in range(G)]
for _ in range(n_records)
]
dataset = GRPODataset()
dataset.storage = type(
"FakeStore",
(),
{
"keys": ["prompts", "responses", "masks", "rewards"],
"_data": {
"prompts": [torch.randint(0, 100, (10,)) for _ in range(n_records)],
"responses": dummy_responses,
"masks": [
[torch.ones(r.shape[0], dtype=torch.int64) for r in resps]
for resps in dummy_responses
],
"rewards": [
torch.rand(G, dtype=torch.float32) for _ in range(n_records)
],
},
},
)()
dataset._build_records()
assert len(dataset) == n_records
for i in range(n_records):
item = dataset[i]
assert len(item["responses"]) == G
assert len(item["masks"]) == G
assert item["rewards"].shape == (G,)
for g in range(G):
assert item["responses"][g].shape == item["masks"][g].shape
+16 -3
View File
@@ -349,7 +349,17 @@ def test_grpo_basic(chat_tokenizer, builder):
assert "responses" in result
assert "masks" in result
assert "rewards" in result
assert len(result["responses"]) == len(result["masks"])
# responses is List[List[int]] — one per response
assert len(result["responses"]) == 4
assert all(isinstance(r, list) for r in result["responses"])
assert all(isinstance(r[0], int) for r in result["responses"])
# masks is List[List[int]] — one per response, matching length
assert len(result["masks"]) == 4
for i in range(4):
assert len(result["masks"][i]) == len(result["responses"][i])
assert result["rewards"] == [1.0, 0.5, 0.8, 0.2]
@@ -362,8 +372,11 @@ def test_grpo_response_tokens_all_trained(chat_tokenizer, builder):
}
result = builder.build(item, config, chat_tokenizer)
masks = result["masks"]
assert all(m == 1 for m in masks)
assert len(masks) == len(result["responses"])
# masks is List[List[int]] — each response's mask should be all 1s
assert len(masks) == 2
for m in masks:
assert all(v == 1 for v in m)
assert len(m) == len(result["responses"][masks.index(m)])
def test_grpo_single_reward(chat_tokenizer, builder):
+67
View File
@@ -7,6 +7,7 @@ from astrai.config.preprocess_config import (
PipelineConfig,
ProcessingConfig,
)
from astrai.preprocessing.packing import PackingStrategyFactory
from astrai.preprocessing.pipeline import Pipeline, filter_by_length
from tests.data.conftest import (
_CHAT_SECTIONS,
@@ -262,3 +263,69 @@ def test_grpo_pipeline(temp_dir, tokenizer_dir):
assert "masks" in meta
assert "rewards" in meta
assert "sequence" not in meta
# ---------------------------------------------------------------------------
# BFD split packing
# ---------------------------------------------------------------------------
_TRU = "keep_start"
def _total_tokens(keys, key="sequence"):
return sum(len(s) for s in keys[key])
def test_bfd_split_preserves_all_tokens():
"""No tokens are lost — split chunks are kept, not truncated away."""
packer = PackingStrategyFactory.create("bfd_split")
max_len = 10
keys = {
"sequence": [list(range(25)), list(range(3))],
"loss_mask": [[1] * 25, [1] * 3],
}
result = packer.apply(keys, max_len, _TRU)
assert _total_tokens(result) == 28
for seq in result["sequence"]:
assert len(seq) <= max_len
def test_bfd_split_chunk_alignment():
"""loss_mask chunks must align with sequence chunks."""
packer = PackingStrategyFactory.create("bfd_split")
max_len = 10
keys = {
"sequence": [list(range(25))],
"loss_mask": [[0] * 5 + [1] * 20],
}
result = packer.apply(keys, max_len, _TRU)
for seq, mask in zip(result["sequence"], result["loss_mask"]):
assert len(seq) == len(mask)
def test_bfd_split_short_unchanged():
"""Sequences under max_packed_len should not be split."""
packer = PackingStrategyFactory.create("bfd_split")
max_len = 10
keys = {"sequence": [list(range(5))], "loss_mask": [[1] * 5]}
result = packer.apply(keys, max_len, _TRU)
assert _total_tokens(result) == 5
assert len(result["sequence"]) >= 1
def test_bfd_split_vs_bfd():
"""bfd loses tokens from over-length sequences; bfd_split does not."""
max_len = 10
keys = {
"sequence": [list(range(25)), list(range(8))],
"loss_mask": [[1] * 25, [1] * 8],
}
bfd = PackingStrategyFactory.create("bfd").apply(keys, max_len, _TRU)
split = PackingStrategyFactory.create("bfd_split").apply(keys, max_len, _TRU)
assert _total_tokens(bfd) < 33
assert _total_tokens(split) == 33
+106
View File
@@ -3,6 +3,7 @@
import torch
from astrai.inference.sample import (
FrequencyPenaltyStrategy,
SamplingPipeline,
TemperatureStrategy,
TopKStrategy,
@@ -125,3 +126,108 @@ def test_module_sample_batch():
assert tokens.shape == (2,)
for t in tokens:
assert 0 <= t < logits.size(-1)
def test_frequency_penalty_noop_when_zero():
logits = torch.tensor([[1.0, 2.0, 3.0]])
input_ids = torch.tensor([[0, 2]])
s = FrequencyPenaltyStrategy(penalty=0.0)
result = s.apply(logits.clone(), input_ids=input_ids)
assert torch.equal(result, logits)
def test_frequency_penalty_noop_when_no_input_ids():
logits = torch.tensor([[1.0, 2.0, 3.0]])
s = FrequencyPenaltyStrategy(penalty=0.5)
result = s.apply(logits.clone())
assert torch.equal(result, logits)
def test_frequency_penalty_single_occurrence():
logits = torch.tensor([[4.0, 1.0, 2.0]])
input_ids = torch.tensor([[0, 2]])
input_mask = torch.tensor([[True, True]])
s = FrequencyPenaltyStrategy(penalty=0.5)
result = s.apply(logits.clone(), input_ids=input_ids, input_mask=input_mask)
assert result[0, 0] == 3.5
assert result[0, 1] == 1.0
assert result[0, 2] == 1.5
def test_frequency_penalty_multiple_occurrences():
logits = torch.tensor([[4.0, 1.0, 2.0]])
input_ids = torch.tensor([[0, 2, 0]])
input_mask = torch.tensor([[True, True, True]])
s = FrequencyPenaltyStrategy(penalty=0.5)
result = s.apply(logits.clone(), input_ids=input_ids, input_mask=input_mask)
assert result[0, 0] == 3.0
assert result[0, 1] == 1.0
assert result[0, 2] == 1.5
def test_frequency_penalty_respects_padding_mask():
logits = torch.tensor([[4.0, 1.0, 2.0]])
input_ids = torch.tensor([[0, 2, 0]])
input_mask = torch.tensor([[True, True, False]])
s = FrequencyPenaltyStrategy(penalty=0.5)
result = s.apply(logits.clone(), input_ids=input_ids, input_mask=input_mask)
assert result[0, 0] == 3.5
assert result[0, 1] == 1.0
assert result[0, 2] == 1.5
def test_frequency_penalty_batch_tensor():
logits = torch.tensor(
[
[4.0, 1.0, 2.0],
[3.0, 5.0, 1.0],
]
)
input_ids = torch.tensor([[0, 2, 0], [1, 1, 0]])
input_mask = torch.tensor([[True, True, True], [True, True, False]])
s = FrequencyPenaltyStrategy(penalty=torch.tensor([0.5, 1.0]))
result = s.apply(logits.clone(), input_ids=input_ids, input_mask=input_mask)
assert result[0, 0] == 3.0
assert result[0, 2] == 1.5
assert result[1, 1] == 3.0
def test_frequency_penalty_negative_penalty_boosts_repeats():
logits = torch.tensor([[4.0, 1.0, 2.0]])
input_ids = torch.tensor([[0, 0]])
input_mask = torch.tensor([[True, True]])
s = FrequencyPenaltyStrategy(penalty=-0.5)
result = s.apply(logits.clone(), input_ids=input_ids, input_mask=input_mask)
assert result[0, 0] == 5.0
def test_frequency_penalty_in_pipeline():
logits = torch.tensor([[5.0, 1.0, 2.0, 3.0]])
input_ids = torch.tensor([[0, 2, 0]])
input_mask = torch.tensor([[True, True, True]])
pipeline = SamplingPipeline(
[
TemperatureStrategy(1.0),
FrequencyPenaltyStrategy(0.5),
]
)
result = pipeline.apply(logits.clone(), input_ids=input_ids, input_mask=input_mask)
assert result[0, 0] == 4.0
assert result[0, 2] == 1.5
def test_sample_with_frequency_penalty():
logits = torch.tensor([[5.0, 1.0, 2.0, 3.0]])
input_ids = torch.tensor([[0, 2, 0]])
input_mask = torch.tensor([[True, True, True]])
tokens = sample(
logits,
temperature=1.0,
top_k=0,
top_p=1.0,
frequency_penalty=0.5,
input_ids=input_ids,
input_mask=input_mask,
)
assert tokens.shape == (1,)
assert 0 <= tokens[0] < logits.size(-1)
+1 -1
View File
@@ -46,7 +46,7 @@ def test_early_stopping_simulation(base_test_env, early_stopping_dataset):
# Resume from latest checkpoint
load_dir = os.path.join(base_test_env["test_dir"], "epoch_0_step_1")
trainer = Trainer(train_config)
trainer.train(resume_dir=load_dir)
trainer.train(param_path=load_dir, resume=True)
# Verify checkpoint was saved at expected step
load_dir = os.path.join(base_test_env["test_dir"], "epoch_1_step_5")
+224
View File
@@ -0,0 +1,224 @@
import pytest
import torch
from astrai.config.model_config import AutoRegressiveLMConfig
from astrai.model.transformer import AutoRegressiveLM
from astrai.trainer.strategy import GRPOStrategy
class _FakeExecutor:
"""Minimal executor stub providing ``unwrap_model`` for ref model creation."""
def unwrap_model(self, model):
return model.state_dict()
def _make_config(vocab_size=200, max_len=64):
return AutoRegressiveLMConfig(
vocab_size=vocab_size,
dim=16,
n_heads=2,
n_kv_heads=1,
dim_ffn=32,
max_len=max_len,
n_layers=2,
norm_eps=1e-5,
)
def _make_model(device):
config = _make_config()
model = AutoRegressiveLM(config).to(device=device)
return model, config
def _make_batch(
batch_size=2, group_size=4, prompt_len=8, response_len=12, device="cpu"
):
"""Construct a GRPO batch with deterministic shapes.
Returns dict with prompts [B, P], responses [B, G, R], masks [B, G, R],
rewards [B, G].
"""
prompts = torch.randint(0, 200, (batch_size, prompt_len), device=device)
responses = torch.randint(
0, 200, (batch_size, group_size, response_len), device=device
)
# All response tokens valid.
masks = torch.ones(batch_size, group_size, response_len, device=device)
# Distinct rewards per group member so std > 0.
rewards = torch.randn(batch_size, group_size, device=device)
return {
"prompts": prompts,
"responses": responses,
"masks": masks,
"rewards": rewards,
}
def _make_frozen_copy(model, device):
"""Create a frozen copy of ``model`` with independent weights loaded."""
config = _make_config()
copy = AutoRegressiveLM(config).to(device=device)
copy.load_state_dict(model.state_dict())
copy.requires_grad_(False)
copy.eval()
return copy
@pytest.fixture
def grpo_strategy():
"""Build a GRPOStrategy with a small real model and fake executor."""
device = "cuda" if torch.cuda.is_available() else "cpu"
model, config = _make_model(device)
old_model = _make_frozen_copy(model, device)
ref_model = _make_frozen_copy(model, device)
strategy = GRPOStrategy(
model=model,
device=device,
old_model=old_model,
ref_model=ref_model,
clip_eps=0.2,
kl_coef=0.01,
group_size=4,
model_fn=lambda c=config: AutoRegressiveLM(c).to(device=device),
executor=_FakeExecutor(),
)
return strategy, device
def test_grpo_loss_is_finite(grpo_strategy):
"""compute_loss returns a finite scalar."""
strategy, device = grpo_strategy
batch = _make_batch(device=device)
loss = strategy.compute_loss(batch)
assert loss.dim() == 0
assert torch.isfinite(loss).item()
def test_grpo_loss_backward(grpo_strategy):
"""Loss is differentiable w.r.t. policy model parameters."""
strategy, device = grpo_strategy
batch = _make_batch(device=device)
loss = strategy.compute_loss(batch)
loss.backward()
# At least some parameter should receive a gradient.
has_grad = any(
p.grad is not None and p.grad.abs().sum().item() > 0
for p in strategy.model.parameters()
)
assert has_grad
def test_grpo_ref_model_not_updated(grpo_strategy):
"""Backward should not populate gradients on ref_model."""
strategy, device = grpo_strategy
batch = _make_batch(device=device)
loss = strategy.compute_loss(batch)
loss.backward()
for p in strategy.ref_model.parameters():
assert p.grad is None
def test_grpo_old_model_not_updated(grpo_strategy):
"""Backward should not populate gradients on old_model."""
strategy, device = grpo_strategy
batch = _make_batch(device=device)
loss = strategy.compute_loss(batch)
loss.backward()
for p in strategy.old_model.parameters():
assert p.grad is None
def test_grpo_prompt_tokens_masked(grpo_strategy):
"""When only prompt-equivalent tokens are unmasked (response mask all 0),
the policy loss should be zero (no valid tokens contribute)."""
strategy, device = grpo_strategy
batch = _make_batch(device=device)
# Zero out all response masks → no response token contributes.
batch["masks"] = torch.zeros_like(batch["masks"])
loss = strategy.compute_loss(batch)
# With no valid tokens, policy_loss term is 0 and KL term is 0.
assert loss.item() == pytest.approx(0.0, abs=1e-6)
def test_grpo_identical_rewards_zero_advantage(grpo_strategy):
"""When all group rewards are identical, advantage is 0 → policy_loss is 0.
Only the KL term remains (which is 0 when policy == ref at init)."""
strategy, device = grpo_strategy
batch = _make_batch(device=device)
batch["rewards"] = torch.ones(batch["rewards"].shape, device=device)
loss = strategy.compute_loss(batch)
# At init policy == old == ref, so ratio == 1, KL == 0; advantage == 0.
assert loss.item() == pytest.approx(0.0, abs=1e-5)
def test_grpo_sync_old_model(grpo_strategy):
"""sync_old_model copies current policy weights into old_model."""
strategy, device = grpo_strategy
# Perturb policy model so it differs from old.
with torch.no_grad():
for p in strategy.model.parameters():
p.add_(0.05)
# old_model should still hold original weights (differ from policy).
policy_sd = strategy.model.state_dict()
old_sd = strategy.old_model.state_dict()
differs_before = any(
not torch.allclose(policy_sd[k], old_sd[k]) for k in policy_sd if k in old_sd
)
assert differs_before
strategy.sync_old_model()
old_sd_after = strategy.old_model.state_dict()
matches = all(
torch.allclose(policy_sd[k], old_sd_after[k])
for k in policy_sd
if k in old_sd_after
)
assert matches
def test_grpo_partial_mask(grpo_strategy):
"""Only the first half of response tokens are valid."""
strategy, device = grpo_strategy
batch = _make_batch(device=device)
B, G, R = batch["masks"].shape
half = R // 2
batch["masks"][:, :, half:] = 0.0
loss = strategy.compute_loss(batch)
assert torch.isfinite(loss).item()
def test_grpo_clipping_effect(grpo_strategy):
"""After diverging policy from ref, ratio should be clipped to [1-eps, 1+eps]
on the surrogate. Verify loss is finite and non-zero for distinct rewards."""
strategy, device = grpo_strategy
# Diverge policy from ref.
with torch.no_grad():
for p in strategy.model.parameters():
p.add_(0.3)
batch = _make_batch(device=device)
loss = strategy.compute_loss(batch)
assert torch.isfinite(loss).item()
# With distinct rewards and diverged policy, loss should be non-trivial.
assert loss.abs().item() > 1e-4
def test_grpo_no_reduction_param():
"""GRPOStrategy.__init__ must not accept ``reduction`` (removed)."""
import inspect
sig = inspect.signature(GRPOStrategy.__init__)
assert "reduction" not in sig.parameters
def test_grpo_shapes_3d_batch(grpo_strategy):
"""Verify compute_loss handles non-square prompt/response lengths."""
strategy, device = grpo_strategy
batch = _make_batch(
batch_size=3, group_size=4, prompt_len=10, response_len=8, device=device
)
loss = strategy.compute_loss(batch)
assert torch.isfinite(loss).item()