Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b99485f462 | ||
|
|
20041d7aa9 | ||
|
|
59248032dc | ||
|
|
ceadc34ea9 | ||
|
|
8ab5631446 |
+1
-1
@@ -1,4 +1,4 @@
|
|||||||
__version__ = "1.3.10"
|
__version__ = "1.3.11"
|
||||||
__author__ = "ViperEkura"
|
__author__ = "ViperEkura"
|
||||||
|
|
||||||
from astrai.config import (
|
from astrai.config import (
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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 // (
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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;
|
||||||
|
|
||||||
|
|||||||
@@ -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++) {
|
||||||
|
|||||||
@@ -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);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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();
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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++) {
|
||||||
|
|||||||
@@ -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;
|
||||||
|
}
|
||||||
+2
-1
@@ -49,4 +49,5 @@ target-version = "py312"
|
|||||||
quote-style = "double"
|
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"]
|
||||||
@@ -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")
|
||||||
|
|
||||||
|
|||||||
@@ -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"]}
|
||||||
|
|||||||
@@ -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():
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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}"
|
||||||
Reference in New Issue
Block a user