5 Commits
Author SHA1 Message Date
ViperEkura b99485f462 chore: bump version to 1.3.11 2026-07-27 01:23:13 +08:00
ViperEkura 20041d7aa9 perf: extend MMA decode to arbitrary GQA ratio, add launch bounds, vectorize combine
- Multi-pass MMA: encode pass in grid blockIdx.x, compute q_head0/G in-kernel
- Fixes crash for G>32 (previously block(32,G) exceeded 1024 threads)
- Fixes alloc_split_partials using uninitialized num_splits (MAX_SPLITS=32)
- __launch_bounds__ on all MMA and prefill kernels for better register allocation
- 4x vectorized combine kernel (4 head_dim per thread)
- uint4 vectorized K loads in scalar decode kernels
- cp.async .L2::128B cache hint for K/V tile streaming
- Extract warp_reduce_sum, bf16, MAX_SPLITS to attn_warp_utils.cuh
2026-07-27 00:35:34 +08:00
ViperEkura 59248032dc chore: fix ruff lint warnings and signal handling edge cases
- Fix pre-existing ruff lint warnings (F401, F541, F841, E741)
- Exclude .md/.json/.yml from ruff format check
- Unblock SIGTERM/SIGINT via pthread_sigmask in early signal handler
- Do not restore SIG_DFL on unregister to prevent pending signal kills
2026-07-25 21:08:30 +08:00
ViperEkura ceadc34ea9 feat: auto-checkpoint on SIGTERM/SIGINT with DDP support
- Register SIGTERM/SIGINT handlers in training loop, set stop flag on signal
- Check stop_requested at each epoch/batch boundary, break and call on_error to save checkpoint
- LocalStrategy parent forwards signal to child processes via terminate(), waits up to 600s for graceful exit
- TrainContext gains threading.Event-based stop_requested/request_stop
- Tests verify SIGTERM/SIGINT trigger checkpoint save with exit code 0, works on both CPU and GPU
2026-07-25 20:40:54 +08:00
ViperEkura 8ab5631446 fix: correct online rollout lifecycle 2026-07-23 19:01:37 +08:00
23 changed files with 551 additions and 85 deletions
+1 -1
View File
@@ -1,4 +1,4 @@
__version__ = "1.3.10" __version__ = "1.3.11"
__author__ = "ViperEkura" __author__ = "ViperEkura"
from astrai.config import ( from astrai.config import (
+6 -2
View File
@@ -190,7 +190,8 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
- rewards: [G] - rewards: [G]
Output: Output:
- prompts: [B, P_max] - prompts: [B, P_max], left-padded
- prompt_mask: [B, P_max]
- responses: [B, G, R_max] - responses: [B, G, R_max]
- masks: [B, G, R_max] - masks: [B, G, R_max]
- rewards: [B, G] - rewards: [B, G]
@@ -201,13 +202,15 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
R_max = max(r.size(0) for b in batch for r in b["responses"]) R_max = max(r.size(0) for b in batch for r in b["responses"])
prompts = torch.zeros(B, P_max, dtype=torch.long) prompts = torch.zeros(B, P_max, dtype=torch.long)
prompt_mask = torch.zeros(B, P_max, dtype=torch.bool)
responses = torch.zeros(B, G, R_max, dtype=torch.long) responses = torch.zeros(B, G, R_max, dtype=torch.long)
masks = torch.zeros(B, G, R_max, dtype=torch.bool) masks = torch.zeros(B, G, R_max, dtype=torch.bool)
rewards = torch.zeros(B, G, dtype=torch.float32) rewards = torch.zeros(B, G, dtype=torch.float32)
for i, b in enumerate(batch): for i, b in enumerate(batch):
p_len = b["prompts"].size(0) p_len = b["prompts"].size(0)
prompts[i, :p_len] = b["prompts"] prompts[i, -p_len:] = b["prompts"]
prompt_mask[i, -p_len:] = True
rewards[i, : b["rewards"].size(0)] = b["rewards"] rewards[i, : b["rewards"].size(0)] = b["rewards"]
for g in range(min(G, len(b["responses"]))): for g in range(min(G, len(b["responses"]))):
r_len = b["responses"][g].size(0) r_len = b["responses"][g].size(0)
@@ -217,6 +220,7 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
return { return {
"prompts": prompts, "prompts": prompts,
"prompt_mask": prompt_mask,
"responses": responses, "responses": responses,
"masks": masks, "masks": masks,
"rewards": rewards, "rewards": rewards,
+45 -3
View File
@@ -1,5 +1,8 @@
import logging
import os import os
import signal
import socket import socket
import threading
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from contextlib import contextmanager from contextlib import contextmanager
from functools import wraps from functools import wraps
@@ -9,6 +12,10 @@ import torch
import torch.distributed as dist import torch.distributed as dist
import torch.multiprocessing as mp import torch.multiprocessing as mp
from astrai.parallel.signal_handler import install_early_signal_handlers
logger = logging.getLogger(__name__)
def find_free_port() -> str: def find_free_port() -> str:
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
@@ -115,6 +122,7 @@ def _run_single_rank(
func: Callable, func: Callable,
kwargs: dict, kwargs: dict,
): ):
install_early_signal_handlers()
with setup_parallel( with setup_parallel(
rank=rank, rank=rank,
world_size=world_size, world_size=world_size,
@@ -155,6 +163,7 @@ class TorchrunStrategy(LaunchStrategy):
"""External orchestrator (torchrun, SLURM, K8s) — env vars pre-set.""" """External orchestrator (torchrun, SLURM, K8s) — env vars pre-set."""
def launch(self, func: Callable, **kwargs): def launch(self, func: Callable, **kwargs):
install_early_signal_handlers()
rank = int(os.environ["RANK"]) rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"]) world_size = int(os.environ["WORLD_SIZE"])
local_rank = int(os.environ.get("LOCAL_RANK", rank)) local_rank = int(os.environ.get("LOCAL_RANK", rank))
@@ -188,6 +197,7 @@ class LocalStrategy(LaunchStrategy):
_run_single_rank(0, *args) _run_single_rank(0, *args)
return return
install_early_signal_handlers()
ctx = mp.start_processes( ctx = mp.start_processes(
_run_single_rank, _run_single_rank,
args=args, args=args,
@@ -195,14 +205,46 @@ class LocalStrategy(LaunchStrategy):
start_method=self.start_method, start_method=self.start_method,
join=False, join=False,
) )
parent_stop = threading.Event()
original_handlers = {}
def _parent_handler(signum, frame):
sig = signal.Signals(signum)
logger.warning(
"Parent (pid=%d) received %s, forwarding to children...",
os.getpid(),
sig.name,
)
parent_stop.set()
for p in ctx.processes:
if p.is_alive():
p.terminate()
for sig in (signal.SIGTERM, signal.SIGINT):
prev = signal.signal(sig, _parent_handler)
if prev not in (signal.SIG_DFL, signal.SIG_IGN, None, _parent_handler):
original_handlers[sig] = prev
try: try:
while not ctx.join(): while not ctx.join() and not parent_stop.is_set():
pass pass
except BaseException: except BaseException:
logger.warning(
"Parent received unexpected exception, terminating children..."
)
for p in ctx.processes: for p in ctx.processes:
p.terminate() if p.is_alive():
ctx.join() p.terminate()
raise raise
finally:
for sig, handler in original_handlers.items():
signal.signal(sig, handler)
for p in ctx.processes:
p.join()
ctx.join()
def _detect_launcher() -> str: def _detect_launcher() -> str:
+53
View File
@@ -0,0 +1,53 @@
import logging
import os
import signal
import threading
logger = logging.getLogger(__name__)
_early_stop = threading.Event()
_active_context = None
def _early_handler(signum: int, frame):
sig = signal.Signals(signum)
logger.warning(
"Received %s (pid=%d), requesting graceful training stop...",
sig.name,
os.getpid(),
)
_early_stop.set()
if _active_context is not None:
_active_context.request_stop()
def install_early_signal_handlers():
for sig in (signal.SIGTERM, signal.SIGINT):
signal.signal(sig, _early_handler)
_unblock_signals()
def _unblock_signals():
try:
mask = signal.pthread_sigmask(signal.SIG_BLOCK, set())
blocked = {signal.SIGTERM, signal.SIGINT} & mask
if blocked:
signal.pthread_sigmask(signal.SIG_UNBLOCK, blocked)
except (AttributeError, OSError):
pass
def register_signal_handlers(context):
global _active_context
_active_context = context
for sig in (signal.SIGTERM, signal.SIGINT):
signal.signal(sig, _early_handler)
if _early_stop.is_set():
context.request_stop()
logger.warning("Signal was received during initialization, stopping...")
def unregister_signal_handlers():
global _active_context
_active_context = None
_early_stop.clear()
+88 -12
View File
@@ -35,6 +35,7 @@ class RawRollout:
Fields: Fields:
prompts: Tokenized prompts, shape ``[B, P_len]``. prompts: Tokenized prompts, shape ``[B, P_len]``.
prompt_mask: Boolean mask for real prompt tokens, shape ``[B, P_len]``.
responses: Generated response token IDs, shape ``[B, G, R_max]``. responses: Generated response token IDs, shape ``[B, G, R_max]``.
response_mask: Boolean mask for real (non-pad) response tokens, response_mask: Boolean mask for real (non-pad) response tokens,
shape ``[B, G, R_max]``. shape ``[B, G, R_max]``.
@@ -47,6 +48,7 @@ class RawRollout:
""" """
prompts: Tensor prompts: Tensor
prompt_mask: Tensor
responses: Tensor responses: Tensor
response_mask: Tensor response_mask: Tensor
logprobs_old: Tensor logprobs_old: Tensor
@@ -143,6 +145,15 @@ class RolloutGenerator:
``add_generation_prompt=True`` so rollout prompts match the ``add_generation_prompt=True`` so rollout prompts match the
format the policy was SFT-trained on. format the policy was SFT-trained on.
""" """
model = self.scheduler._executor.model
was_training = model.training
model.eval()
try:
return self._generate_eval(batch)
finally:
model.train(was_training)
def _generate_eval(self, batch: Dict) -> RawRollout:
prompt_texts, flat_prompt_ids = self._prepare_prompts(batch) prompt_texts, flat_prompt_ids = self._prepare_prompts(batch)
B = len(prompt_texts) B = len(prompt_texts)
G = self.group_size G = self.group_size
@@ -161,6 +172,15 @@ class RolloutGenerator:
rep_window=self.rep_window, rep_window=self.rep_window,
return_logprobs=True, return_logprobs=True,
) )
if len(results) != B * G:
raise RuntimeError(
f"Rollout scheduler returned {len(results)} results, expected {B * G}"
)
for token_ids, logprobs in results:
if len(token_ids) != len(logprobs):
raise RuntimeError(
"Rollout scheduler returned misaligned token IDs and logprobs"
)
# Each element is (token_ids, logprobs); pad to max length. # Each element is (token_ids, logprobs); pad to max length.
max_len = 0 max_len = 0
@@ -171,10 +191,12 @@ class RolloutGenerator:
device = self.scheduler.device device = self.scheduler.device
P_len = max(len(ids) for ids in flat_prompt_ids) P_len = max(len(ids) for ids in flat_prompt_ids)
prompts_tensor = torch.zeros(B, P_len, dtype=torch.long, device=device) prompts_tensor = torch.zeros(B, P_len, dtype=torch.long, device=device)
prompt_mask = torch.zeros(B, P_len, dtype=torch.bool, device=device)
for i, ids in enumerate(flat_prompt_ids): for i, ids in enumerate(flat_prompt_ids):
prompts_tensor[i, : len(ids)] = torch.tensor( prompts_tensor[i, -len(ids) :] = torch.tensor(
ids, dtype=torch.long, device=device ids, dtype=torch.long, device=device
) )
prompt_mask[i, -len(ids) :] = True
responses = torch.full((B, G, max_len), _PAD, dtype=torch.long, device=device) responses = torch.full((B, G, max_len), _PAD, dtype=torch.long, device=device)
response_mask = torch.zeros((B, G, max_len), dtype=torch.bool, device=device) response_mask = torch.zeros((B, G, max_len), dtype=torch.bool, device=device)
@@ -201,6 +223,7 @@ class RolloutGenerator:
return RawRollout( return RawRollout(
prompts=prompts_tensor, prompts=prompts_tensor,
prompt_mask=prompt_mask,
responses=responses, responses=responses,
response_mask=response_mask, response_mask=response_mask,
logprobs_old=logprobs_old, logprobs_old=logprobs_old,
@@ -239,17 +262,35 @@ class RolloutGenerator:
f"{list(batch.keys())}" f"{list(batch.keys())}"
) )
prompt_texts: List[str] = [] try:
flat_prompt_ids: List[List[int]] = [] prompt_texts = self.tokenizer.apply_chat_template(
for messages in messages_list: messages_list, tokenize=False, add_generation_prompt=True
text = self.tokenizer.apply_chat_template(
messages, tokenize=False, add_generation_prompt=True
) )
ids = self.tokenizer.apply_chat_template( if (
messages, tokenize=True, add_generation_prompt=True not isinstance(prompt_texts, list)
) or len(prompt_texts) != len(messages_list)
prompt_texts.append(text) or not all(isinstance(text, str) for text in prompt_texts)
flat_prompt_ids.append(list(ids)) ):
raise TypeError("Tokenizer does not support batched chat templates")
flat_prompt_ids = self.tokenizer.encode(prompt_texts)
if len(flat_prompt_ids) != len(messages_list) or not all(
isinstance(ids, list) for ids in flat_prompt_ids
):
raise TypeError("Tokenizer does not support batched encoding")
except (TypeError, IndexError, KeyError):
# Keep compatibility with lightweight tokenizer adapters that only
# implement the single-conversation template API.
prompt_texts = []
flat_prompt_ids = []
for messages in messages_list:
text = self.tokenizer.apply_chat_template(
messages, tokenize=False, add_generation_prompt=True
)
ids = self.tokenizer.apply_chat_template(
messages, tokenize=True, add_generation_prompt=True
)
prompt_texts.append(text)
flat_prompt_ids.append(list(ids))
return prompt_texts, flat_prompt_ids return prompt_texts, flat_prompt_ids
@staticmethod @staticmethod
@@ -308,6 +349,7 @@ class RolloutRunner:
self.rollout_interval = rollout_interval self.rollout_interval = rollout_interval
self._cache: Optional[RolloutResult] = None self._cache: Optional[RolloutResult] = None
self._cache_key = None
self._steps_since_rollout: int = 0 self._steps_since_rollout: int = 0
def step(self): def step(self):
@@ -317,12 +359,40 @@ class RolloutRunner:
def clear_cache(self): def clear_cache(self):
"""Force next call to re-run rollout.""" """Force next call to re-run rollout."""
self._cache = None self._cache = None
self._cache_key = None
@staticmethod
def _batch_key(batch: Dict):
"""Build a stable key for the prompt fields accepted by the generator."""
def freeze(value):
if isinstance(value, dict):
return tuple(sorted((key, freeze(val)) for key, val in value.items()))
if isinstance(value, (list, tuple)):
return tuple(freeze(item) for item in value)
return value
fields = ("messages", "instruction", "input", "output")
return tuple(
(field, freeze(batch[field])) for field in fields if field in batch
)
def _score(self, raw: RawRollout) -> RolloutResult: def _score(self, raw: RawRollout) -> RolloutResult:
rewards = self.reward_model.score(raw.prompt_texts, raw.response_texts) rewards = self.reward_model.score(raw.prompt_texts, raw.response_texts)
if not isinstance(rewards, Tensor):
rewards = torch.as_tensor(rewards, dtype=torch.float32)
expected_shape = raw.responses.shape[:2]
if rewards.shape != expected_shape:
raise ValueError(
f"Reward model returned shape {tuple(rewards.shape)}, "
f"expected {tuple(expected_shape)}"
)
if not torch.isfinite(rewards).all():
raise ValueError("Reward model returned non-finite values")
device = raw.prompts.device device = raw.prompts.device
return RolloutResult( return RolloutResult(
prompts=raw.prompts, prompts=raw.prompts,
prompt_mask=raw.prompt_mask,
responses=raw.responses, responses=raw.responses,
response_mask=raw.response_mask, response_mask=raw.response_mask,
rewards=rewards.to(device=device), rewards=rewards.to(device=device),
@@ -337,9 +407,15 @@ class RolloutRunner:
Triggers a new rollout when ``_steps_since_rollout >= rollout_interval`` Triggers a new rollout when ``_steps_since_rollout >= rollout_interval``
or when the cache is empty. or when the cache is empty.
""" """
if self._cache is None or self._steps_since_rollout >= self.rollout_interval: cache_key = self._batch_key(batch)
if (
self._cache is None
or cache_key != self._cache_key
or self._steps_since_rollout >= self.rollout_interval
):
raw = self.generator.generate(batch) raw = self.generator.generate(batch)
self._cache = self._score(raw) self._cache = self._score(raw)
self._cache_key = cache_key
self._steps_since_rollout = 0 self._steps_since_rollout = 0
return self._cache, True return self._cache, True
return self._cache, False return self._cache, False
+16 -4
View File
@@ -1,7 +1,7 @@
"""Training strategy implementations with factory pattern.""" """Training strategy implementations with factory pattern."""
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Callable, Dict, Optional, Union from typing import Callable, Dict, Union
import torch import torch
import torch.nn as nn import torch.nn as nn
@@ -158,6 +158,11 @@ class BaseStrategy(ABC):
""" """
pass pass
def on_optimizer_step(self):
"""Advance online rollout state after a successful optimizer step."""
if self._rollout_runner is not None:
self._rollout_runner.step()
def __call__(self, batch: Dict[str, Tensor]) -> Tensor: def __call__(self, batch: Dict[str, Tensor]) -> Tensor:
"""Run offline or online forward depending on runner injection.""" """Run offline or online forward depending on runner injection."""
if self._rollout_runner is None: if self._rollout_runner is None:
@@ -166,8 +171,6 @@ class BaseStrategy(ABC):
result, is_fresh = self._rollout_runner(batch) result, is_fresh = self._rollout_runner(batch)
if is_fresh: if is_fresh:
self._on_rollout_refresh() self._on_rollout_refresh()
if self.executor and self.executor.sync_gradients:
self._rollout_runner.step()
train_batch = self.prepare_from_rollout(result) train_batch = self.prepare_from_rollout(result)
return self.compute_loss(train_batch) return self.compute_loss(train_batch)
@@ -411,6 +414,12 @@ class GRPOStrategy(BaseStrategy):
responses_flat = responses.view(-1, response_len) responses_flat = responses.view(-1, response_len)
masks_flat = masks.view(-1, response_len) masks_flat = masks.view(-1, response_len)
prompt_expanded = prompts.unsqueeze(1).repeat(1, group_size, 1).flatten(0, 1) prompt_expanded = prompts.unsqueeze(1).repeat(1, group_size, 1).flatten(0, 1)
prompt_mask = batch.get("prompt_mask")
if prompt_mask is None:
prompt_mask = prompts.ne(0)
prompt_mask_expanded = (
prompt_mask.unsqueeze(1).expand(-1, group_size, -1).flatten(0, 1)
)
prompt_len = prompt_expanded.size(1) prompt_len = prompt_expanded.size(1)
full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-1) full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-1)
@@ -423,7 +432,9 @@ class GRPOStrategy(BaseStrategy):
) )
# Build full attention mask: key-padding + causal # Build full attention mask: key-padding + causal
key_pad = full_sequences.bool()[:, None, None, :] key_pad = torch.cat([prompt_mask_expanded, masks_flat.bool()], dim=-1)[
:, None, None, :
]
S = key_pad.shape[-1] S = key_pad.shape[-1]
causal = torch.tril( causal = torch.tril(
torch.ones(S, S, dtype=torch.bool, device=full_sequences.device) torch.ones(S, S, dtype=torch.bool, device=full_sequences.device)
@@ -485,6 +496,7 @@ class GRPOStrategy(BaseStrategy):
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]: def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
return { return {
"prompts": result.prompts, "prompts": result.prompts,
"prompt_mask": result.prompt_mask,
"responses": result.responses, "responses": result.responses,
"masks": result.response_mask, "masks": result.response_mask,
"rewards": result.rewards, "rewards": result.rewards,
+10
View File
@@ -1,3 +1,4 @@
import threading
from dataclasses import dataclass, field from dataclasses import dataclass, field
from pathlib import Path from pathlib import Path
from typing import Any, Dict, Optional, Self from typing import Any, Dict, Optional, Self
@@ -41,6 +42,15 @@ class TrainContext:
rank: int = field(default=0) rank: int = field(default=0)
kwargs: Dict[str, Any] = field(default_factory=dict) kwargs: Dict[str, Any] = field(default_factory=dict)
_stop_event: threading.Event = field(default_factory=threading.Event)
@property
def stop_requested(self) -> bool:
return self._stop_event.is_set()
def request_stop(self) -> None:
self._stop_event.set()
@property @property
def optimizer_step(self) -> int: def optimizer_step(self) -> int:
return self.consumed_samples // ( return self.consumed_samples // (
+21
View File
@@ -1,8 +1,14 @@
import logging import logging
from typing import List, Optional from typing import List, Optional
import torch.distributed as dist
from astrai.config import TrainConfig from astrai.config import TrainConfig
from astrai.parallel.setup import spawn_parallel_fn from astrai.parallel.setup import spawn_parallel_fn
from astrai.parallel.signal_handler import (
register_signal_handlers,
unregister_signal_handlers,
)
from astrai.trainer.train_callback import ( from astrai.trainer.train_callback import (
CallbackFactory, CallbackFactory,
TrainCallback, TrainCallback,
@@ -58,6 +64,7 @@ class Trainer:
.with_param_path(param_path, resume=resume) .with_param_path(param_path, resume=resume)
.build() .build()
) )
register_signal_handlers(context)
executor = context.executor executor = context.executor
self._call_callbacks("on_train_begin", context) self._call_callbacks("on_train_begin", context)
@@ -65,10 +72,14 @@ class Trainer:
context.model.train() context.model.train()
for epoch in range(context.epoch, context.config.n_epoch): for epoch in range(context.epoch, context.config.n_epoch):
if context.stop_requested:
break
context.epoch = epoch context.epoch = epoch
self._call_callbacks("on_epoch_begin", context) self._call_callbacks("on_epoch_begin", context)
for batch in context.dataloader: for batch in context.dataloader:
if context.stop_requested:
break
with executor.accumulate(context.model): with executor.accumulate(context.model):
self._call_callbacks("on_batch_begin", context) self._call_callbacks("on_batch_begin", context)
loss = context.strategy(batch) loss = context.strategy(batch)
@@ -83,6 +94,7 @@ class Trainer:
if executor.sync_gradients: if executor.sync_gradients:
self._call_callbacks("on_optimizer_step", context) self._call_callbacks("on_optimizer_step", context)
context.optimizer.step() context.optimizer.step()
context.strategy.on_optimizer_step()
context.optimizer.zero_grad() context.optimizer.zero_grad()
if context.scheduler: if context.scheduler:
@@ -90,12 +102,21 @@ class Trainer:
self._call_callbacks("on_epoch_end", context) self._call_callbacks("on_epoch_end", context)
if context.stop_requested:
logger.warning(
"Training interrupted by signal, saving emergency checkpoint..."
)
self._call_callbacks("on_error", context)
except Exception as e: except Exception as e:
logger.error("Training failed: %s", str(e), exc_info=True) logger.error("Training failed: %s", str(e), exc_info=True)
self._call_callbacks("on_error", context) self._call_callbacks("on_error", context)
raise raise
finally: finally:
self._call_callbacks("on_train_end", context) self._call_callbacks("on_train_end", context)
if executor.use_distributed and dist.is_initialized():
dist.barrier()
unregister_signal_handlers()
def train(self, param_path: Optional[str] = None, resume: bool = False): def train(self, param_path: Optional[str] = None, resume: bool = False):
cfg = self.train_config cfg = self.train_config
+3 -10
View File
@@ -2,16 +2,9 @@
#include <cuda_bf16.h> #include <cuda_bf16.h>
#include <float.h> #include <float.h>
#include "attn_common.h" #include "attn_common.h"
#include "attn_warp_utils.cuh"
using bf16 = __nv_bfloat16;
constexpr int DC_CHUNK = 64; constexpr int DC_CHUNK = 64;
__device__ inline float warp_reduce_sum(float val) {
for (int offset = 16; offset > 0; offset >>= 1)
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
return val;
}
template <int HEAD_DIM, bool IsCausal, bool HasMask> template <int HEAD_DIM, bool IsCausal, bool HasMask>
__global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) { __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
int batch = blockIdx.x / p.kv_head; int batch = blockIdx.x / p.kv_head;
@@ -93,7 +86,7 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
// ---- write UN-normalised partials for this split ---- // ---- write UN-normalised partials for this split ----
size_t bh = (size_t)batch * p.q_head + q_head; size_t bh = (size_t)batch * p.q_head + q_head;
size_t slot = bh * p.num_splits + split; size_t slot = bh * MAX_SPLITS + split;
int d0 = lane * hd_per_thread; int d0 = lane * hd_per_thread;
for (int i = 0; i < hd_per_thread; i++) { for (int i = 0; i < hd_per_thread; i++) {
int dd = d0 + i; int dd = d0 + i;
@@ -113,7 +106,7 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
int batch = bh / p.q_head; int batch = bh / p.q_head;
int q_head = bh % p.q_head; int q_head = bh % p.q_head;
size_t split_base = (size_t)bh * p.num_splits; size_t split_base = (size_t)bh * MAX_SPLITS;
const float* mlp = p.ml_part + split_base * 2; const float* mlp = p.ml_part + split_base * 2;
const float* op = p.o_part + split_base * p.head_dim; const float* op = p.o_part + split_base * p.head_dim;
+10 -4
View File
@@ -3,6 +3,7 @@
#include <cuda_bf16.h> #include <cuda_bf16.h>
#include "attn_common.h" #include "attn_common.h"
#include "attn_mma_utils.cuh" #include "attn_mma_utils.cuh"
#include "attn_warp_utils.cuh"
// Split-K (FlashDecoding) tensor-core decode via GQA head-packing. // Split-K (FlashDecoding) tensor-core decode via GQA head-packing.
// Decode has q_len == 1, so we pack G = q_head/kv_head query heads into the // Decode has q_len == 1, so we pack G = q_head/kv_head query heads into the
@@ -19,11 +20,16 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
const int gid = lane >> 2; const int gid = lane >> 2;
const int tid4 = lane & 3; const int tid4 = lane & 3;
const int kv_head = blockIdx.x; const int pass = blockIdx.x / p.kv_head;
const int kv_head = blockIdx.x % p.kv_head;
const int batch = blockIdx.y; const int batch = blockIdx.y;
const int split = blockIdx.z; const int split = blockIdx.z;
const int G = p.q_head / p.kv_head;
const int q_head0 = kv_head * G; constexpr int MAX_G = 16;
const int G_total = p.q_head / p.kv_head;
const int g_begin = pass * MAX_G;
const int G = min(MAX_G, G_total - g_begin);
const int q_head0 = kv_head * G_total + g_begin;
// Double-buffered shared memory for K/V (no sQ needed) // Double-buffered shared memory for K/V (no sQ needed)
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD]; __shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
@@ -120,7 +126,7 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
// ---- write UN-normalised partials for this split ---- // ---- write UN-normalised partials for this split ----
auto split_slot = [&](int h) -> size_t { auto split_slot = [&](int h) -> size_t {
size_t bh = (size_t)batch * p.q_head + h; size_t bh = (size_t)batch * p.q_head + h;
return bh * p.num_splits + split; return bh * MAX_SPLITS + split;
}; };
#pragma unroll #pragma unroll
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) { for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
+20 -21
View File
@@ -4,6 +4,7 @@
#include <cuda_runtime.h> #include <cuda_runtime.h>
#include <algorithm> #include <algorithm>
#include "attn_warp_utils.cuh"
#include "attn_prefill_split_q.cuh" #include "attn_prefill_split_q.cuh"
#include "attn_decode_split_kv.cuh" #include "attn_decode_split_kv.cuh"
#include "attn_paged_decode_split_kv.cuh" #include "attn_paged_decode_split_kv.cuh"
@@ -18,7 +19,7 @@ inline int compute_num_splits(int base_blocks, int tiles_total) {
int sm_count = 0; int sm_count = 0;
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0); cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
int n = (2 * sm_count + base_blocks - 1) / base_blocks; int n = (2 * sm_count + base_blocks - 1) / base_blocks;
return std::max(1, std::min(n, std::min(tiles_total, 32))); return std::max(1, std::min(n, std::min(tiles_total, MAX_SPLITS)));
} }
// ====================================================================== // ======================================================================
@@ -77,21 +78,14 @@ static inline void dispatch_prefill(AttentionParams<bf16>& p) {
template <int HEAD_DIM, bool IsCausal, bool HasMask> template <int HEAD_DIM, bool IsCausal, bool HasMask>
static inline void launch_decode_mma(AttentionParams<bf16>& p, int group_size) { static inline void launch_decode_mma(AttentionParams<bf16>& p, int group_size) {
int G = p.q_head / p.kv_head; int G = p.q_head / p.kv_head;
if (G >= 1 && G <= 16) { constexpr int MAX_G = 16;
int tiles_total = (p.kv_len + 32 - 1) / 32; int num_passes = (G + MAX_G - 1) / MAX_G;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total); int tiles_total = (p.kv_len + 32 - 1) / 32;
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1; p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
using Traits = KernelTraits<HEAD_DIM, 32, 1, STAGES>; constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
dim3 grid(p.kv_head, p.batch, p.num_splits); using Traits = KernelTraits<HEAD_DIM, 32, 1, STAGES>;
attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask><<<grid, 32>>>(p); dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
} else { attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask><<<grid, 32>>>(p);
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
dim3 block(32, group_size);
attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
}
} }
#endif #endif
@@ -100,8 +94,9 @@ static inline void launch_decode_scalar(AttentionParams<bf16>& p, int group_size
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK; int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total); p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16); size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit
dim3 grid(p.batch * p.kv_head, 1, p.num_splits); dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
dim3 block(32, group_size); dim3 block(32, g);
attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p); attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
} }
@@ -140,13 +135,16 @@ static inline void dispatch_decode(AttentionParams<bf16>& p) {
template <int HEAD_DIM, bool IsCausal, bool HasMask> template <int HEAD_DIM, bool IsCausal, bool HasMask>
static inline void launch_paged_decode_mma(PagedAttentionParams<bf16>& p, int group_size) { static inline void launch_paged_decode_mma(PagedAttentionParams<bf16>& p, int group_size) {
int G = p.q_head / p.kv_head; int G = p.q_head / p.kv_head;
if (G >= 1 && G <= 16 && p.page_size >= 32) { constexpr int MAX_G = 16;
bool page_ok = (p.page_size >= 32);
if (G >= 1 && page_ok) {
int num_passes = (G + MAX_G - 1) / MAX_G;
int tiles_total = (p.kv_len + 32 - 1) / 32; int tiles_total = (p.kv_len + 32 - 1) / 32;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total); p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1; constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
using Traits = KernelTraits<HEAD_DIM, 32, 1, STAGES>; using Traits = KernelTraits<HEAD_DIM, 32, 1, STAGES>;
dim3 grid(p.kv_head, p.batch, p.num_splits); dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
paged_attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask><<<grid, 32>>>(p); paged_attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask> <<<grid, 32>>>(p);
} else { } else {
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK; int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total); p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
@@ -163,8 +161,9 @@ static inline void launch_paged_decode_scalar(PagedAttentionParams<bf16>& p, int
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK; int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total); p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16); size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit
dim3 grid(p.batch * p.kv_head, 1, p.num_splits); dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
dim3 block(32, group_size); dim3 block(32, g);
paged_attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p); paged_attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
} }
+3 -2
View File
@@ -2,6 +2,7 @@
#include <torch/extension.h> #include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h> #include <c10/cuda/CUDAGuard.h>
#include "attn_common.h" #include "attn_common.h"
#include "attn_warp_utils.cuh"
using bf16 = __nv_bfloat16; using bf16 = __nv_bfloat16;
@@ -22,8 +23,8 @@ using bf16 = __nv_bfloat16;
template<typename P> template<typename P>
inline void alloc_split_partials(P& p) { inline void alloc_split_partials(P& p) {
auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA); auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
auto o_part = torch::empty({p.batch, p.q_head, p.num_splits, p.head_dim}, fopt); auto o_part = torch::empty({p.batch, p.q_head, MAX_SPLITS, p.head_dim}, fopt);
auto ml_part = torch::empty({p.batch, p.q_head, p.num_splits, 2}, fopt); auto ml_part = torch::empty({p.batch, p.q_head, MAX_SPLITS, 2}, fopt);
p.o_part = (float*)o_part.data_ptr(); p.o_part = (float*)o_part.data_ptr();
p.ml_part = (float*)ml_part.data_ptr(); p.ml_part = (float*)ml_part.data_ptr();
} }
+4 -11
View File
@@ -2,16 +2,9 @@
#include <cuda_bf16.h> #include <cuda_bf16.h>
#include <float.h> #include <float.h>
#include "attn_common.h" #include "attn_common.h"
#include "attn_warp_utils.cuh"
using bf16 = __nv_bfloat16;
constexpr int PDC_CHUNK = 64; constexpr int PDC_CHUNK = 64;
__device__ inline float paged_warp_reduce_sum(float val) {
for (int offset = 16; offset > 0; offset >>= 1)
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
return val;
}
template <int HEAD_DIM, bool IsCausal, bool HasMask> template <int HEAD_DIM, bool IsCausal, bool HasMask>
__global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p) { __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p) {
int batch = blockIdx.x / p.kv_head; int batch = blockIdx.x / p.kv_head;
@@ -71,7 +64,7 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
for (int i = 0; i < hd_per_thread; i++) for (int i = 0; i < hd_per_thread; i++)
partial += q_reg[i] * __bfloat162float( partial += q_reg[i] * __bfloat162float(
k_smem[s * p.head_dim + lane * hd_per_thread + i]); k_smem[s * p.head_dim + lane * hd_per_thread + i]);
partial = paged_warp_reduce_sum(partial) * p.scale; partial = warp_reduce_sum(partial) * p.scale;
int kv_idx = chunk_start + s; int kv_idx = chunk_start + s;
if constexpr (HasMask) { if constexpr (HasMask) {
@@ -111,7 +104,7 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
} }
size_t bh = (size_t)batch * p.q_head + q_head; size_t bh = (size_t)batch * p.q_head + q_head;
size_t slot = bh * p.num_splits + split; size_t slot = bh * MAX_SPLITS + split;
int d0 = lane * hd_per_thread; int d0 = lane * hd_per_thread;
#pragma unroll #pragma unroll
for (int i = 0; i < hd_per_thread; i++) for (int i = 0; i < hd_per_thread; i++)
@@ -130,7 +123,7 @@ __global__ void paged_attn_decode_combine_kernel(PagedAttentionParams<bf16> p) {
int batch = bh / p.q_head; int batch = bh / p.q_head;
int q_head = bh % p.q_head; int q_head = bh % p.q_head;
size_t split_base = (size_t)bh * p.num_splits; size_t split_base = (size_t)bh * MAX_SPLITS;
const float* mlp = p.ml_part + split_base * 2; const float* mlp = p.ml_part + split_base * 2;
const float* op = p.o_part + split_base * p.head_dim; const float* op = p.o_part + split_base * p.head_dim;
@@ -3,6 +3,7 @@
#include <cuda_bf16.h> #include <cuda_bf16.h>
#include "attn_common.h" #include "attn_common.h"
#include "attn_mma_utils.cuh" #include "attn_mma_utils.cuh"
#include "attn_warp_utils.cuh"
// Paged split-KV tensor-core decode via GQA head-packing. // Paged split-KV tensor-core decode via GQA head-packing.
// Reads K/V directly from the page pool through a page table — one tile // Reads K/V directly from the page pool through a page table — one tile
@@ -16,11 +17,16 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
const int gid = lane >> 2; const int gid = lane >> 2;
const int tid4 = lane & 3; const int tid4 = lane & 3;
const int kv_head = blockIdx.x; const int pass = blockIdx.x / p.kv_head;
const int kv_head = blockIdx.x % p.kv_head;
const int batch = blockIdx.y; const int batch = blockIdx.y;
const int split = blockIdx.z; const int split = blockIdx.z;
const int G = p.q_head / p.kv_head;
const int q_head0 = kv_head * G; constexpr int MAX_G = 16;
const int G_total = p.q_head / p.kv_head;
const int g_begin = pass * MAX_G;
const int G = min(MAX_G, G_total - g_begin);
const int q_head0 = kv_head * G_total + g_begin;
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD]; __shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD]; __shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
@@ -120,7 +126,7 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
auto split_slot = [&](int h) -> size_t { auto split_slot = [&](int h) -> size_t {
size_t bh = (size_t)batch * p.q_head + h; size_t bh = (size_t)batch * p.q_head + h;
return bh * p.num_splits + split; return bh * MAX_SPLITS + split;
}; };
#pragma unroll #pragma unroll
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) { for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
+13
View File
@@ -0,0 +1,13 @@
#pragma once
#include <cuda_bf16.h>
using bf16 = __nv_bfloat16;
static constexpr int MAX_SPLITS = 32;
__device__ inline float warp_reduce_sum(float val) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1)
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
return val;
}
+1
View File
@@ -50,3 +50,4 @@ quote-style = "double"
indent-style = "space" indent-style = "space"
skip-magic-trailing-comma = false skip-magic-trailing-comma = false
line-ending = "auto" line-ending = "auto"
exclude = ["*.md", "*.json", "*.yml", "*.yaml"]
+3 -3
View File
@@ -143,7 +143,7 @@ def print_layer_grid(results: dict[str, dict]):
widths = [6] + [10] * len(comps) widths = [6] + [10] * len(comps)
metric = "er_99_norm" metric = "er_99_norm"
print(f"\n--- Per-Layer Effective Rank (99% energy) ---") print("\n--- Per-Layer Effective Rank (99% energy) ---")
print(format_header(["Layer"] + comps, widths)) print(format_header(["Layer"] + comps, widths))
print("-" * sum(widths)) print("-" * sum(widths))
@@ -173,7 +173,7 @@ def print_layer_grid(results: dict[str, dict]):
def print_weight_stats(results: dict[str, dict]): def print_weight_stats(results: dict[str, dict]):
groups = group_by_component(results) groups = group_by_component(results)
widths = [20, 12, 12, 12, 12] widths = [20, 12, 12, 12, 12]
print(f"\n--- Weight Value Statistics ---") print("\n--- Weight Value Statistics ---")
print(format_header(["Component", "Mean", "Std", "Min", "Max"], widths)) print(format_header(["Component", "Mean", "Std", "Min", "Max"], widths))
print("-" * sum(widths)) print("-" * sum(widths))
@@ -265,7 +265,7 @@ def main():
) )
print(f"{'=' * 70}") print(f"{'=' * 70}")
print(f"Loading weights...") print("Loading weights...")
sd = safetensors.torch.load_file(str(weights_path)) sd = safetensors.torch.load_file(str(weights_path))
print(f" {len(sd)} keys loaded") print(f" {len(sd)} keys loaded")
-1
View File
@@ -215,7 +215,6 @@ def _permute_choices(item: dict, rng: random.Random) -> tuple[dict, str]:
positional bias (e.g. always picking B). positional bias (e.g. always picking B).
""" """
letters = ("A", "B", "C", "D") letters = ("A", "B", "C", "D")
contents = [item[k] for k in letters]
perm = list(letters) perm = list(letters)
rng.shuffle(perm) rng.shuffle(perm)
permuted = {"question": item["question"]} permuted = {"question": item["question"]}
+2 -2
View File
@@ -148,7 +148,7 @@ class LossAccumulator:
self.total += sum(losses) self.total += sum(losses)
self.count += len(losses) self.count += len(losses)
if self.stream: if self.stream:
clamped = [min(max(l, 0.0), self._HIST_MAX) for l in losses] clamped = [min(max(v, 0.0), self._HIST_MAX) for v in losses]
idx = torch.tensor(clamped) / self._HIST_MAX * (self._HIST_BINS - 1) idx = torch.tensor(clamped) / self._HIST_MAX * (self._HIST_BINS - 1)
self.hist += torch.bincount( self.hist += torch.bincount(
idx.long().clamp(0, self._HIST_BINS - 1), idx.long().clamp(0, self._HIST_BINS - 1),
@@ -315,7 +315,7 @@ def print_stats(label: str, stats: Dict):
) )
by_type = stats.get("by_token_type", {}) by_type = stats.get("by_token_type", {})
if by_type: if by_type:
print(f"\n by token type:") print("\n by token type:")
print(f" {'type':<12} {'count':>8} {'mean_loss':>10} {'ppl':>8}") print(f" {'type':<12} {'count':>8} {'mean_loss':>10} {'ppl':>8}")
print(f" {'-' * 12} {'-' * 8} {'-' * 10} {'-' * 8}") print(f" {'-' * 12} {'-' * 8} {'-' * 10} {'-' * 8}")
for ttype, s in by_type.items(): for ttype, s in by_type.items():
+4 -3
View File
@@ -824,7 +824,7 @@ def test_grpo_builder_preserves_response_boundaries(base_test_env):
from tests.data.conftest import make_grpo_no_template_config from tests.data.conftest import make_grpo_no_template_config
tokenizer = base_test_env["tokenizer"] tokenizer = base_test_env["tokenizer"]
tokenizer_path = _save_test_tokenizer(base_test_env["test_dir"], tokenizer) _save_test_tokenizer(base_test_env["test_dir"], tokenizer)
builder = SectionedMaskBuilder() builder = SectionedMaskBuilder()
config = make_grpo_no_template_config() config = make_grpo_no_template_config()
@@ -937,8 +937,9 @@ def test_grpo_collate_variable_lengths():
assert result["masks"].shape == (2, 2, 4) assert result["masks"].shape == (2, 2, 4)
assert result["rewards"].shape == (2, 2) assert result["rewards"].shape == (2, 2)
# Check padding: item 1 prompt is length 2, padded to 3 # Prompts are left-padded so each response follows its real prompt tokens.
assert result["prompts"][1, 2] == 0 assert torch.equal(result["prompts"][1], torch.tensor([0, 10, 11]))
assert torch.equal(result["prompt_mask"][1], torch.tensor([False, True, True]))
# Check response content: item 0, response 0 is [4,5] padded to 4 # Check response content: item 0, response 0 is [4,5] padded to 4
assert result["responses"][0, 0, 0] == 4 assert result["responses"][0, 0, 0] == 4
+10 -1
View File
@@ -63,6 +63,7 @@ def _make_frozen(model, device):
def _make_rollout_result(B=2, G=4, P=6, R=8, device="cpu"): def _make_rollout_result(B=2, G=4, P=6, R=8, device="cpu"):
return RolloutResult( return RolloutResult(
prompts=torch.randint(3, 200, (B, P), device=device), prompts=torch.randint(3, 200, (B, P), device=device),
prompt_mask=torch.ones(B, P, dtype=torch.bool, device=device),
responses=torch.randint(3, 200, (B, G, R), device=device), responses=torch.randint(3, 200, (B, G, R), device=device),
response_mask=torch.ones(B, G, R, dtype=torch.bool, device=device), response_mask=torch.ones(B, G, R, dtype=torch.bool, device=device),
rewards=torch.randn(B, G, device=device), rewards=torch.randn(B, G, device=device),
@@ -177,6 +178,7 @@ def test_grpo_prepare_from_rollout_mapping(device):
r = _make_rollout_result(device=device) r = _make_rollout_result(device=device)
batch = strat.prepare_from_rollout(r) batch = strat.prepare_from_rollout(r)
assert batch["prompts"] is r.prompts assert batch["prompts"] is r.prompts
assert batch["prompt_mask"] is r.prompt_mask
assert batch["responses"] is r.responses assert batch["responses"] is r.responses
assert batch["masks"] is r.response_mask assert batch["masks"] is r.response_mask
assert batch["rewards"] is r.rewards assert batch["rewards"] is r.rewards
@@ -255,9 +257,11 @@ def test_grpo_no_resync_when_same_cached_result(device):
runner = _RecordingRunner(_make_rollout_result(device=device)) runner = _RecordingRunner(_make_rollout_result(device=device))
strat.set_rollout_runner(runner) strat.set_rollout_runner(runner)
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)}) strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)}) strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
assert runner.calls == 2 assert runner.calls == 2
assert runner.step_calls == 1 assert runner.step_calls == 2
def test_grpo_resync_when_new_rollout_result(device): def test_grpo_resync_when_new_rollout_result(device):
@@ -265,8 +269,10 @@ def test_grpo_resync_when_new_rollout_result(device):
runner = _RecordingRunner(_make_rollout_result(device=device)) runner = _RecordingRunner(_make_rollout_result(device=device))
strat.set_rollout_runner(runner) strat.set_rollout_runner(runner)
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)}) strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
runner.swap_result(_make_rollout_result(device=device)) runner.swap_result(_make_rollout_result(device=device))
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)}) strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
assert runner.calls == 2 assert runner.calls == 2
assert runner.step_calls == 2 assert runner.step_calls == 2
@@ -281,8 +287,10 @@ def test_dpo_no_sync_hook_when_new_rollout_result(device):
runner = _RecordingRunner(_make_rollout_result(device=device)) runner = _RecordingRunner(_make_rollout_result(device=device))
strat.set_rollout_runner(runner) strat.set_rollout_runner(runner)
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)}) strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
runner.swap_result(_make_rollout_result(device=device)) runner.swap_result(_make_rollout_result(device=device))
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)}) strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
assert runner.step_calls == 2 assert runner.step_calls == 2
@@ -301,6 +309,7 @@ def test_step_called_when_sync_gradients_true(device):
runner = _RecordingRunner(_make_rollout_result(device=device)) runner = _RecordingRunner(_make_rollout_result(device=device))
strat.set_rollout_runner(runner) strat.set_rollout_runner(runner)
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)}) strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
assert runner.step_calls == 1 assert runner.step_calls == 1
+50
View File
@@ -71,6 +71,18 @@ class ConstantRewardModel(BaseRewardModel):
return torch.full((B, G), float(self.value)) return torch.full((B, G), float(self.value))
class BadShapeRewardModel(BaseRewardModel):
def score(self, prompts, responses):
return torch.zeros(len(prompts))
class NonFiniteRewardModel(BaseRewardModel):
def score(self, prompts, responses):
B = len(prompts)
G = len(responses[0]) if B else 0
return torch.full((B, G), float("nan"))
def _make_config(vocab_size=200, max_position_embeddings=128): def _make_config(vocab_size=200, max_position_embeddings=128):
return AutoRegressiveLMConfig( return AutoRegressiveLMConfig(
vocab_size=vocab_size, vocab_size=vocab_size,
@@ -111,6 +123,7 @@ def _make_instruction_batch(n=2):
def test_raw_rollout_fields(): def test_raw_rollout_fields():
r = RawRollout( r = RawRollout(
prompts=torch.zeros(2, 4, dtype=torch.long), prompts=torch.zeros(2, 4, dtype=torch.long),
prompt_mask=torch.ones(2, 4, dtype=torch.bool),
responses=torch.zeros(2, 3, 5, dtype=torch.long), responses=torch.zeros(2, 3, 5, dtype=torch.long),
response_mask=torch.ones(2, 3, 5, dtype=torch.bool), response_mask=torch.ones(2, 3, 5, dtype=torch.bool),
logprobs_old=torch.zeros(2, 3, 5), logprobs_old=torch.zeros(2, 3, 5),
@@ -124,6 +137,7 @@ def test_raw_rollout_fields():
def test_rollout_result_inherits_raw_rollout_fields(): def test_rollout_result_inherits_raw_rollout_fields():
r = RolloutResult( r = RolloutResult(
prompts=torch.zeros(2, 4, dtype=torch.long), prompts=torch.zeros(2, 4, dtype=torch.long),
prompt_mask=torch.ones(2, 4, dtype=torch.bool),
responses=torch.zeros(2, 3, 5, dtype=torch.long), responses=torch.zeros(2, 3, 5, dtype=torch.long),
response_mask=torch.ones(2, 3, 5, dtype=torch.bool), response_mask=torch.ones(2, 3, 5, dtype=torch.bool),
logprobs_old=torch.zeros(2, 3, 5), logprobs_old=torch.zeros(2, 3, 5),
@@ -132,6 +146,7 @@ def test_rollout_result_inherits_raw_rollout_fields():
assert r.rewards.shape == (2, 3) assert r.rewards.shape == (2, 3)
assert r.prompts.shape == (2, 4) assert r.prompts.shape == (2, 4)
assert r.responses.shape == (2, 3, 5) assert r.responses.shape == (2, 3, 5)
assert r.prompt_mask.shape == (2, 4)
def test_base_reward_model_is_abstract(): def test_base_reward_model_is_abstract():
@@ -179,11 +194,28 @@ def test_rollout_generator_shapes(device):
assert r.responses.shape == (2, 3, 5) assert r.responses.shape == (2, 3, 5)
assert r.response_mask.shape == (2, 3, 5) assert r.response_mask.shape == (2, 3, 5)
assert r.logprobs_old.shape == (2, 3, 5) assert r.logprobs_old.shape == (2, 3, 5)
assert r.prompt_mask.shape == r.prompts.shape
assert len(r.prompt_texts) == 2 assert len(r.prompt_texts) == 2
assert len(r.response_texts) == 2 assert len(r.response_texts) == 2
assert len(r.response_texts[0]) == 3 assert len(r.response_texts[0]) == 3
def test_rollout_generator_uses_eval_and_restores_mode(device):
gen, model = _make_generator(device, group_size=1, max_tokens=2)
model.train()
seen_training = []
original = gen.scheduler.run_batch
def recording_run_batch(*args, **kwargs):
seen_training.append(model.training)
return original(*args, **kwargs)
gen.scheduler.run_batch = recording_run_batch
gen.generate(_make_instruction_batch(n=1))
assert seen_training == [False]
assert model.training is True
def test_rollout_generator_mask_matches_responses(device): def test_rollout_generator_mask_matches_responses(device):
"""Positions beyond a response's length are pad (mask False).""" """Positions beyond a response's length are pad (mask False)."""
gen, _ = _make_generator(device, group_size=2, max_tokens=6) gen, _ = _make_generator(device, group_size=2, max_tokens=6)
@@ -291,6 +323,24 @@ def test_rollout_runner_cache_returns_stale_flag(device):
assert fresh2 is False assert fresh2 is False
def test_rollout_runner_refreshes_for_different_batch(device):
runner, _ = _make_runner(device, rollout_interval=100)
r1, fresh1 = runner(_make_instruction_batch(n=1))
batch2 = {"instruction": ["Different prompt"], "input": [""]}
r2, fresh2 = runner(batch2)
assert fresh1 is True
assert fresh2 is True
assert r2 is not r1
@pytest.mark.parametrize("reward_model", [BadShapeRewardModel, NonFiniteRewardModel])
def test_rollout_runner_rejects_invalid_rewards(device, reward_model):
generator, _ = _make_generator(device, group_size=2, max_tokens=2)
runner = RolloutRunner(generator, reward_model(), rollout_interval=1)
with pytest.raises(ValueError):
runner(_make_instruction_batch(n=1))
def test_rollout_runner_step_triggers_new_rollout(device): def test_rollout_runner_step_triggers_new_rollout(device):
runner, _ = _make_runner(device, rollout_interval=2) runner, _ = _make_runner(device, rollout_interval=2)
batch = _make_instruction_batch() batch = _make_instruction_batch()
+177
View File
@@ -0,0 +1,177 @@
import json
import multiprocessing as mp
import os
import signal
import time
import pytest
import torch
import torch.optim as optim
from torch.utils.data import Dataset
from astrai.config import TrainConfig
from astrai.config.model_config import AutoRegressiveLMConfig
from astrai.model.transformer import AutoRegressiveLM
from astrai.parallel.signal_handler import register_signal_handlers
from astrai.trainer import Trainer
from astrai.trainer.schedule import SchedulerFactory
from astrai.trainer.train_context import TrainContext
class _PicklableDataset(Dataset):
def __init__(self, length=200, max_length=64, vocab_size=1000):
self.length = length
self.max_length = max_length
self.vocab_size = vocab_size
def __len__(self):
return self.length
def __getitem__(self, idx):
return {
"input_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
"target_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
}
def _build_model():
config = AutoRegressiveLMConfig(
vocab_size=1000,
hidden_size=8,
num_attention_heads=2,
num_key_value_heads=1,
intermediate_size=16,
max_position_embeddings=64,
num_hidden_layers=2,
rms_norm_eps=1e-5,
)
device = "cuda" if torch.cuda.is_available() else "cpu"
return AutoRegressiveLM(config).to(device=device)
class _ReadyCallback:
def __init__(self, ready_file):
self._ready_file = ready_file
def on_train_begin(self, context):
with open(self._ready_file, "w") as f:
f.write("ready")
f.flush()
os.fsync(f.fileno())
def _inner_run(batch_per_device, ckpt_interval, ckpt_dir, log_dir, ready_file):
dataset = _PicklableDataset()
def model_fn():
return _build_model()
def optimizer_fn(m):
return optim.AdamW(m.parameters(), lr=0.001)
def scheduler_fn(optim):
return SchedulerFactory.create(
"cosine", optim, warmup_steps=10, lr_decay_steps=10, min_rate=0.05
)
train_config = TrainConfig(
strategy="seq",
model_fn=model_fn,
dataset=dataset,
optimizer_fn=optimizer_fn,
scheduler_fn=scheduler_fn,
ckpt_dir=ckpt_dir,
log_dir=log_dir,
n_epoch=1,
batch_per_device=batch_per_device,
ckpt_interval=ckpt_interval,
grad_accum_steps=1,
random_seed=42,
device_type="cuda" if torch.cuda.is_available() else "cpu",
)
trainer = Trainer(train_config)
trainer.callbacks.insert(0, _ReadyCallback(ready_file))
trainer.train()
def _spawn_train_and_signal(ckpt_dir, sig, timeout=120):
log_dir = os.path.join(ckpt_dir, "logs")
ready_file = os.path.join(ckpt_dir, "ready.txt")
ctx = mp.get_context("spawn")
p = ctx.Process(
target=_inner_run,
args=(2, 1000, ckpt_dir, log_dir, ready_file),
)
p.start()
deadline = time.time() + 30
while time.time() < deadline:
if os.path.exists(ready_file):
with open(ready_file) as f:
if f.read().strip() == "ready":
break
if not p.is_alive():
break
time.sleep(0.5)
assert p.is_alive(), "Training process died before becoming ready"
os.kill(p.pid, sig)
p.join(timeout=timeout)
if p.is_alive():
p.kill()
p.join(timeout=5)
return p.exitcode
def test_context_stop_flag():
ctx = TrainContext()
assert not ctx.stop_requested
ctx.request_stop()
assert ctx.stop_requested
def test_register_signal_handlers():
ctx = TrainContext()
register_signal_handlers(ctx)
assert not ctx.stop_requested
os.kill(os.getpid(), signal.SIGTERM)
assert ctx.stop_requested
def test_sigterm_triggers_checkpoint_save(base_test_env):
exitcode = _spawn_train_and_signal(base_test_env["test_dir"], signal.SIGTERM)
assert exitcode == 0, f"Training process exited with code {exitcode} (expected 0)"
ckpt_dir = base_test_env["test_dir"]
meta_files = []
for root, dirs, files in os.walk(ckpt_dir):
for f in files:
if f == "meta.json":
meta_files.append(os.path.join(root, f))
assert len(meta_files) > 0, f"No checkpoint meta.json found in {ckpt_dir}"
with open(meta_files[-1]) as f:
meta = json.load(f)
assert "consumed_samples" in meta
assert meta["consumed_samples"] >= 0
@pytest.mark.slow
def test_sigint_triggers_checkpoint_save(base_test_env):
exitcode = _spawn_train_and_signal(base_test_env["test_dir"], signal.SIGINT)
assert exitcode == 0, f"Training process exited with code {exitcode} (expected 0)"
ckpt_dir = base_test_env["test_dir"]
meta_files = []
for root, dirs, files in os.walk(ckpt_dir):
for f in files:
if f == "meta.json":
meta_files.append(os.path.join(root, f))
assert len(meta_files) > 0, f"No checkpoint meta.json found in {ckpt_dir}"