Compare commits
38
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
738cb8f128 | ||
|
|
28d1bd07cf | ||
|
|
02625739fe | ||
|
|
f688cd9c5a | ||
|
|
8055027df7 | ||
|
|
3067a8e1a6 | ||
|
|
97114b95a4 | ||
|
|
32fd03a025 | ||
|
|
21bf37dd83 | ||
|
|
5b67d5865a | ||
|
|
df979b4469 | ||
|
|
deb2d7e127 | ||
|
|
fc47319240 | ||
|
|
22cf798d81 | ||
|
|
164be9708b | ||
|
|
6a97524db4 | ||
|
|
c8b1e40f71 | ||
|
|
bcaa2d1ae0 | ||
|
|
8206afefd9 | ||
|
|
646b1b0f46 | ||
|
|
8150ab6c32 | ||
|
|
0b0693a0a2 | ||
|
|
115192c67c | ||
|
|
c2b04d8458 | ||
|
|
db487ab48b | ||
|
|
a95794d3db | ||
|
|
39f84f3b4c | ||
|
|
9f7cf50c56 | ||
|
|
d9a0c72149 | ||
|
|
5ab18bec48 | ||
|
|
2e29ed45d3 | ||
|
|
5ba21f4eb3 | ||
|
|
c26a47b0df | ||
|
|
b1a87b22bb | ||
|
|
07625057f2 | ||
|
|
53c804e233 | ||
|
|
05c7432964 | ||
|
|
4de42d83c2 |
+1
-1
@@ -4,6 +4,6 @@
|
||||
# Allow necessary files
|
||||
!astrai/
|
||||
!scripts/
|
||||
!assets/
|
||||
!docs/
|
||||
!pyproject.toml
|
||||
!README.md
|
||||
|
||||
+1
-1
@@ -24,7 +24,7 @@
|
||||
!/.dockerignore
|
||||
!/Dockerfile
|
||||
!/docker-compose.yml
|
||||
!/assets/**
|
||||
!/docs/**
|
||||
!/CONTRIBUTING.md
|
||||
!/LICENSE
|
||||
!/pyproject.toml
|
||||
|
||||
+1
-1
@@ -43,7 +43,7 @@ ENV PATH="/opt/venv/bin:$PATH"
|
||||
# Copy application code
|
||||
COPY astrai/ ./astrai/
|
||||
COPY scripts/ ./scripts/
|
||||
COPY assets/ ./assets/
|
||||
COPY docs/ ./docs/
|
||||
COPY pyproject.toml .
|
||||
COPY README.md .
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
<div align="center">
|
||||
|
||||
<img src="assets/images/logo.png" width="auto" alt="Logo">
|
||||
<img src="docs/images/logo.png" width="auto" alt="Logo">
|
||||
<p>
|
||||
<strong>A lightweight Transformer training & inference framework</strong>
|
||||
</p>
|
||||
@@ -17,7 +17,7 @@
|
||||
|
||||
<div align="center">
|
||||
<a href="#english">English</a> •
|
||||
<a href="assets/docs/README-zh-CN.md">中文</a> •
|
||||
<a href="docs/README-zh-CN.md">中文</a> •
|
||||
<a href="https://github.com/ViperEkura/AstrAI/issues">Issue Tracker</a> •
|
||||
<a href="https://github.com/ViperEkura/AstrAI/discussions">Discussions</a> •
|
||||
<a href="https://huggingface.co/ViperEkura">HuggingFace</a>
|
||||
@@ -213,18 +213,23 @@ curl -X POST http://localhost:8000/v1/messages \
|
||||
curl http://localhost:8000/health
|
||||
```
|
||||
|
||||
See [Inference Guide](assets/docs/inference.md) for SSE streaming format, error codes, and stats endpoint.
|
||||
See [Inference Guide](docs/guides/inference.md) for SSE streaming format, error codes, and stats endpoint.
|
||||
|
||||
### Documentation
|
||||
|
||||
| Document | Description |
|
||||
|----------|-------------|
|
||||
| [CLI Reference](./assets/docs/params.md) | Parameters for all CLI tools (train, server, generate, preprocess) |
|
||||
| [Architecture](./assets/docs/architecture.md) | System architecture, class diagram & design patterns |
|
||||
| [Training](./assets/docs/training.md) | Training loop, strategies & formulas |
|
||||
| [Inference](./assets/docs/inference.md) | KVCache, continuous batching, sampling & HTTP API |
|
||||
| [Data Flow](./assets/docs/dataflow.md) | Data pipeline, storage backends & dataset architecture |
|
||||
| [Preprocessing](./assets/docs/preprocessing.md) | Declarative JSON-driven data preprocessing |
|
||||
| [Get Started](./docs/get-started.md) | Installation and quickstart |
|
||||
| [CLI Reference](./docs/guides/params.md) | Parameters for all CLI tools (train, server, generate, preprocess) |
|
||||
| [Preprocessing](./docs/guides/preprocessing.md) | Declarative JSON-driven data preprocessing |
|
||||
| [Training](./docs/guides/training.md) | Training loop, strategies & formulas |
|
||||
| [Inference](./docs/guides/inference.md) | KVCache, continuous batching, sampling & HTTP API |
|
||||
| [Evaluation](./docs/guides/evaluation.md) | HumanEval, MMLU, PPL, ROUGE, IFD, IFEval |
|
||||
| [Distributed](./docs/guides/distributed.md) | Multi-GPU DDP / FSDP training |
|
||||
| [Architecture](./docs/developer/architecture.md) | System architecture, class diagram & design patterns |
|
||||
| [Data Flow](./docs/developer/dataflow.md) | Data pipeline, storage backends & dataset architecture |
|
||||
| [Internals](./docs/developer/internals.md) | Training internals: loss formulas, callback lifecycle, KV cache |
|
||||
| [CUDA Kernels](./docs/developer/cuda_kernels.md) | Custom CUDA attention kernels & benchmarks |
|
||||
|
||||
### Contributing
|
||||
|
||||
|
||||
@@ -1,6 +1,9 @@
|
||||
__version__ = "1.3.11"
|
||||
__author__ = "ViperEkura"
|
||||
|
||||
import logging
|
||||
import os
|
||||
|
||||
from astrai.config import (
|
||||
AutoRegressiveLMConfig,
|
||||
BaseModelConfig,
|
||||
@@ -53,6 +56,30 @@ from astrai.trainer import (
|
||||
Trainer,
|
||||
)
|
||||
|
||||
|
||||
def setup_logging(level: str = "INFO"):
|
||||
"""Attach a handler to the ``astrai`` logger (only, not root).
|
||||
|
||||
Call once per process, e.g. at the top of CLI scripts.
|
||||
Set ``ASTR_LOG_LEVEL`` to override the default ``INFO``.
|
||||
"""
|
||||
_logger = logging.getLogger("astrai")
|
||||
if _logger.handlers:
|
||||
return
|
||||
_level = getattr(
|
||||
logging, os.environ.get("ASTR_LOG_LEVEL", level).upper(), logging.INFO
|
||||
)
|
||||
_logger.setLevel(_level)
|
||||
_handler = logging.StreamHandler()
|
||||
_handler.setFormatter(
|
||||
logging.Formatter(
|
||||
"%(asctime)s | %(levelname)-7s | %(name)s | %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S",
|
||||
)
|
||||
)
|
||||
_logger.addHandler(_handler)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AutoRegressiveLM",
|
||||
"AutoRegressiveLMConfig",
|
||||
@@ -94,5 +121,6 @@ __all__ = [
|
||||
"only_on_rank",
|
||||
"run_server",
|
||||
"sample",
|
||||
"setup_logging",
|
||||
"spawn_parallel_fn",
|
||||
]
|
||||
|
||||
+20
-80
@@ -1,92 +1,32 @@
|
||||
import json
|
||||
from dataclasses import MISSING, dataclass, fields
|
||||
from dataclasses import asdict
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Optional, Self, Union, get_type_hints
|
||||
from typing import Any, Dict, Self, Union
|
||||
|
||||
from pydantic import ConfigDict
|
||||
from pydantic.dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass
|
||||
@dataclass(config=ConfigDict(use_attribute_docstrings=True))
|
||||
class BaseConfig:
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
d = {}
|
||||
for fld in fields(self):
|
||||
v = getattr(self, fld.name)
|
||||
if isinstance(v, (str, int, float, bool)):
|
||||
d[fld.name] = v
|
||||
elif v is None:
|
||||
d[fld.name] = None
|
||||
elif isinstance(v, (dict, list, tuple)):
|
||||
try:
|
||||
val = list(v) if isinstance(v, tuple) else v
|
||||
json.dumps(val)
|
||||
d[fld.name] = val
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
elif isinstance(v, BaseConfig):
|
||||
d[fld.name] = v.to_dict()
|
||||
elif hasattr(v, "__dataclass_fields__"):
|
||||
sub = {}
|
||||
for f in fields(v):
|
||||
a = getattr(v, f.name)
|
||||
sub[f.name] = list(a) if isinstance(a, tuple) else a
|
||||
d[fld.name] = sub
|
||||
return d
|
||||
result = {}
|
||||
for k, v in asdict(self).items():
|
||||
if isinstance(v, tuple):
|
||||
v = list(v)
|
||||
try:
|
||||
json.dumps(v)
|
||||
result[k] = v
|
||||
except (TypeError, ValueError):
|
||||
# Skip non-serializable runtime objects (e.g. model_fn, dataset).
|
||||
# TrainConfig mixes hyperparams with callables/datasets; only the
|
||||
# JSON-serializable subset is written to checkpoint meta.
|
||||
pass
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, d: Dict[str, Any]) -> Self:
|
||||
hints = get_type_hints(cls)
|
||||
inst = cls.__new__(cls)
|
||||
for fld in fields(cls):
|
||||
if fld.name in d:
|
||||
v = d[fld.name]
|
||||
target = cls._unwrap_optional(hints.get(fld.name))
|
||||
if target is not None:
|
||||
try:
|
||||
v = cls._coerce(v, target)
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
object.__setattr__(inst, fld.name, v)
|
||||
elif fld.default is not MISSING:
|
||||
object.__setattr__(inst, fld.name, fld.default)
|
||||
elif fld.default_factory is not MISSING:
|
||||
object.__setattr__(inst, fld.name, fld.default_factory())
|
||||
else:
|
||||
object.__setattr__(inst, fld.name, None)
|
||||
return inst
|
||||
|
||||
@staticmethod
|
||||
def _unwrap_optional(tp) -> Optional[type]:
|
||||
if tp is None:
|
||||
return None
|
||||
origin = getattr(tp, "__origin__", None)
|
||||
if origin is not None:
|
||||
args = getattr(tp, "__args__", ())
|
||||
non_none = [a for a in args if a is not type(None)]
|
||||
return non_none[0] if non_none else None
|
||||
return tp
|
||||
|
||||
@staticmethod
|
||||
def _coerce(value: Any, target_type: type) -> Any:
|
||||
if target_type is bool and isinstance(value, bool):
|
||||
return value
|
||||
if (
|
||||
target_type is int
|
||||
and isinstance(value, (int, float))
|
||||
and not isinstance(value, bool)
|
||||
):
|
||||
return int(value)
|
||||
if (
|
||||
target_type is float
|
||||
and isinstance(value, (int, float))
|
||||
and not isinstance(value, bool)
|
||||
):
|
||||
return float(value)
|
||||
if target_type is str and isinstance(value, str):
|
||||
return value
|
||||
if isinstance(value, target_type):
|
||||
return value
|
||||
if isinstance(value, dict) and issubclass(target_type, BaseConfig):
|
||||
return target_type.from_dict(value)
|
||||
raise TypeError
|
||||
return cls(**d)
|
||||
|
||||
@classmethod
|
||||
def from_file(cls, path: Union[str, Path]) -> Self:
|
||||
|
||||
@@ -1,9 +1,14 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from pydantic import field_validator
|
||||
from pydantic.dataclasses import dataclass
|
||||
|
||||
from astrai.config.base import BaseConfig
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
_ATTN_TYPES = frozenset({"gqa", "mla"})
|
||||
_FFN_TYPES = frozenset({"mlp", "moe"})
|
||||
|
||||
|
||||
class ConfigFactory(BaseFactory[BaseConfig]):
|
||||
"""Factory that dispatches config classes by ``model_type``."""
|
||||
@@ -17,7 +22,12 @@ class ConfigFactory(BaseFactory[BaseConfig]):
|
||||
|
||||
@dataclass
|
||||
class BaseModelConfig(BaseConfig):
|
||||
"""Base config with ``model_type`` dispatch and file I/O."""
|
||||
"""Base config with ``model_type`` dispatch and file I/O.
|
||||
|
||||
Args:
|
||||
model_type (Optional[str]): Model type identifier for AutoModel dispatch. Defaults to None.
|
||||
neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0.
|
||||
"""
|
||||
|
||||
model_type: Optional[str] = None
|
||||
neftune_alpha: float = 0.0
|
||||
@@ -26,7 +36,34 @@ class BaseModelConfig(BaseConfig):
|
||||
@dataclass
|
||||
@ConfigFactory.register("autoregressive_lm")
|
||||
class AutoRegressiveLMConfig(BaseModelConfig):
|
||||
"""Configuration for autoregressive language model."""
|
||||
"""Configuration for autoregressive language model.
|
||||
|
||||
Args:
|
||||
model_type (Optional[str]): Model type identifier for AutoModel dispatch. Defaults to None.
|
||||
neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0.
|
||||
vocab_size (Optional[int]): Vocabulary size. Defaults to None.
|
||||
hidden_size (Optional[int]): Hidden dimension size. Defaults to None.
|
||||
num_hidden_layers (Optional[int]): Number of transformer layers. Defaults to None.
|
||||
rms_norm_eps (Optional[float]): Epsilon for RMSNorm. Defaults to None.
|
||||
intermediate_size (Optional[int]): Intermediate size in FFN. Defaults to None.
|
||||
tie_word_embeddings (Optional[bool]): Whether to tie embedding and lm_head weights. Defaults to None.
|
||||
max_position_embeddings (Optional[int]): Maximum sequence length the model was trained with. Defaults to None.
|
||||
rope_theta (Optional[float]): Base frequency for RoPE. Defaults to None.
|
||||
rope_scaling (Optional[dict]): RoPE scaling config, e.g. {"type": "linear", "factor": 4.0}. Defaults to None.
|
||||
attn_type (str): Attention type: 'gqa' or 'mla'. Defaults to "gqa".
|
||||
num_attention_heads (Optional[int]): Number of query attention heads. Defaults to None.
|
||||
num_key_value_heads (Optional[int]): Number of key/value heads for GQA. Defaults to None.
|
||||
use_qk_norm (Optional[bool]): Whether to apply RMSNorm to Q/K. Defaults to None.
|
||||
use_gated_attention (Optional[bool]): Whether to use gated attention. Defaults to None.
|
||||
kv_lora_rank (Optional[int]): KV compression rank, MLA only. Defaults to None.
|
||||
qk_nope_head_dim (Optional[int]): Non-RoPE head dimension, MLA only. Defaults to None.
|
||||
qk_rope_head_dim (Optional[int]): RoPE head dimension, MLA only. Defaults to None.
|
||||
ffn_type (str): FFN type: 'mlp' or 'moe'. Defaults to "mlp".
|
||||
n_routed_experts (Optional[int]): Number of routed experts, MoE only. Defaults to None.
|
||||
n_shared_experts (Optional[int]): Number of shared experts, MoE only. Defaults to None.
|
||||
n_activated_experts (Optional[int]): Number of activated experts per token, MoE only. Defaults to None.
|
||||
topk_method (Optional[str]): Top-k routing method, MoE only. Defaults to None.
|
||||
"""
|
||||
|
||||
vocab_size: Optional[int] = None
|
||||
hidden_size: Optional[int] = None
|
||||
@@ -34,49 +71,91 @@ class AutoRegressiveLMConfig(BaseModelConfig):
|
||||
rms_norm_eps: Optional[float] = None
|
||||
intermediate_size: Optional[int] = None
|
||||
tie_word_embeddings: Optional[bool] = None
|
||||
|
||||
max_position_embeddings: Optional[int] = None
|
||||
rope_theta: Optional[float] = None
|
||||
rope_scaling: Optional[dict] = None
|
||||
|
||||
attn_type: str = "gqa"
|
||||
num_attention_heads: Optional[int] = None
|
||||
num_key_value_heads: Optional[int] = None
|
||||
use_qk_norm: Optional[bool] = None
|
||||
use_gated_attention: Optional[bool] = None
|
||||
|
||||
kv_lora_rank: Optional[int] = None
|
||||
qk_nope_head_dim: Optional[int] = None
|
||||
qk_rope_head_dim: Optional[int] = None
|
||||
|
||||
ffn_type: str = "mlp"
|
||||
n_routed_experts: Optional[int] = None
|
||||
n_shared_experts: Optional[int] = None
|
||||
n_activated_experts: Optional[int] = None
|
||||
topk_method: Optional[str] = None
|
||||
|
||||
@field_validator("attn_type")
|
||||
def _validate_attn_type(cls, v: str) -> str:
|
||||
if v not in _ATTN_TYPES:
|
||||
raise ValueError(
|
||||
f"attn_type must be one of {sorted(_ATTN_TYPES)}, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
@field_validator("ffn_type")
|
||||
def _validate_ffn_type(cls, v: str) -> str:
|
||||
if v not in _FFN_TYPES:
|
||||
raise ValueError(f"ffn_type must be one of {sorted(_FFN_TYPES)}, got {v!r}")
|
||||
return v
|
||||
|
||||
|
||||
@dataclass
|
||||
@ConfigFactory.register("embedding")
|
||||
class EncoderConfig(BaseModelConfig):
|
||||
"""Configuration for embedding encoder model."""
|
||||
"""Configuration for embedding encoder model.
|
||||
|
||||
Args:
|
||||
model_type (Optional[str]): Model type identifier for AutoModel dispatch. Defaults to None.
|
||||
neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0.
|
||||
vocab_size (Optional[int]): Vocabulary size. Defaults to None.
|
||||
hidden_size (Optional[int]): Hidden dimension size. Defaults to None.
|
||||
num_hidden_layers (Optional[int]): Number of transformer layers. Defaults to None.
|
||||
rms_norm_eps (Optional[float]): Epsilon for RMSNorm. Defaults to None.
|
||||
intermediate_size (Optional[int]): Intermediate size in FFN. Defaults to None.
|
||||
max_position_embeddings (Optional[int]): Maximum sequence length the model was trained with. Defaults to None.
|
||||
rope_theta (Optional[float]): Base frequency for RoPE. Defaults to None.
|
||||
rope_scaling (Optional[dict]): RoPE scaling config, e.g. {"type": "linear", "factor": 4.0}. Defaults to None.
|
||||
attn_type (str): Attention type: 'gqa' or 'mla'. Defaults to "gqa".
|
||||
num_attention_heads (Optional[int]): Number of query attention heads. Defaults to None.
|
||||
num_key_value_heads (Optional[int]): Number of key/value heads for GQA. Defaults to None.
|
||||
use_qk_norm (Optional[bool]): Whether to apply RMSNorm to Q/K. Defaults to None.
|
||||
use_gated_attention (Optional[bool]): Whether to use gated attention. Defaults to None.
|
||||
ffn_type (str): FFN type: 'mlp' or 'moe'. Defaults to "mlp".
|
||||
pooling_type (Optional[str]): Pooling strategy for embedding, e.g. 'mean', 'cls'. Defaults to None.
|
||||
normalize_embeddings (Optional[bool]): Whether to L2-normalize output embeddings. Defaults to None.
|
||||
"""
|
||||
|
||||
vocab_size: Optional[int] = None
|
||||
hidden_size: Optional[int] = None
|
||||
num_hidden_layers: Optional[int] = None
|
||||
rms_norm_eps: Optional[float] = None
|
||||
intermediate_size: Optional[int] = None
|
||||
|
||||
max_position_embeddings: Optional[int] = None
|
||||
rope_theta: Optional[float] = None
|
||||
rope_scaling: Optional[dict] = None
|
||||
|
||||
attn_type: str = "gqa"
|
||||
num_attention_heads: Optional[int] = None
|
||||
num_key_value_heads: Optional[int] = None
|
||||
use_qk_norm: Optional[bool] = None
|
||||
use_gated_attention: Optional[bool] = None
|
||||
|
||||
ffn_type: str = "mlp"
|
||||
pooling_type: Optional[str] = None
|
||||
normalize_embeddings: Optional[bool] = None
|
||||
|
||||
@field_validator("attn_type")
|
||||
def _validate_attn_type(cls, v: str) -> str:
|
||||
if v not in _ATTN_TYPES:
|
||||
raise ValueError(
|
||||
f"attn_type must be one of {sorted(_ATTN_TYPES)}, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
@field_validator("ffn_type")
|
||||
def _validate_ffn_type(cls, v: str) -> str:
|
||||
if v not in _FFN_TYPES:
|
||||
raise ValueError(f"ffn_type must be one of {sorted(_FFN_TYPES)}, got {v!r}")
|
||||
return v
|
||||
|
||||
@@ -5,11 +5,19 @@ modes, both driven declaratively through ``input.sections`` or
|
||||
``input.sources``.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from dataclasses import field
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from pydantic import field_validator
|
||||
from pydantic.dataclasses import dataclass
|
||||
|
||||
from astrai.config.base import BaseConfig
|
||||
|
||||
_PACKING_STRATEGIES = frozenset({"simple", "bfd", "bfd_split"})
|
||||
_TRUNCATION_MODES = frozenset({"keep_start", "keep_end"})
|
||||
_STORAGE_FORMATS = frozenset({"bin", "jsonl"})
|
||||
_POSITION_IDS_MODES = frozenset({"none", "doc_reset", "continuous"})
|
||||
|
||||
|
||||
@dataclass
|
||||
class InputConfig(BaseConfig):
|
||||
@@ -25,6 +33,10 @@ class InputConfig(BaseConfig):
|
||||
"chosen": {"sections": [{"field": "chosen", ...}]},
|
||||
"rejected": {"sections": [{"field": "rejected", ...}]},
|
||||
}}}
|
||||
|
||||
Args:
|
||||
sections (Optional[List[Dict]]): Section list for single-output mode. Defaults to None.
|
||||
sources (Optional[Dict[str, Dict]]): Source map for multi-output mode, DPO/GRPO. Defaults to None.
|
||||
"""
|
||||
|
||||
sections: Optional[List[Dict]] = None
|
||||
@@ -33,34 +45,17 @@ class InputConfig(BaseConfig):
|
||||
|
||||
@dataclass
|
||||
class ProcessingConfig(BaseConfig):
|
||||
"""Processing configuration.
|
||||
"""Processing configuration for tokenization and packing.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
max_seq_len : int
|
||||
Maximum sequence length (default: 2048).
|
||||
min_chars : int
|
||||
Minimum number of characters to keep (default: 50).
|
||||
max_chars : int
|
||||
Maximum number of characters to keep (default: 2_000_000).
|
||||
max_items : Optional[int]
|
||||
Maximum number of items to process (default: None, unlimited).
|
||||
batch_size : int
|
||||
Number of records tokenized together (default: 256).
|
||||
packing_strategy : str
|
||||
How to pack sequences into a contiguous stream.
|
||||
|
||||
- ``"simple"``: sequential concatenation (default, backward compatible).
|
||||
- ``"bfd"``: best-fit decreasing bin packing, minimises wasted tokens.
|
||||
- ``"bfd_split"``: BFD with over-length sequences split into chunks.
|
||||
max_packed_len : int
|
||||
Maximum length of a packed bin. Sequences longer than this are
|
||||
truncated or split depending on ``packing_strategy`` (default: 8192).
|
||||
truncation_mode : str
|
||||
How to truncate sequences longer than ``max_packed_len``.
|
||||
|
||||
- ``"keep_start"``: keep the first ``max_packed_len`` tokens (default).
|
||||
- ``"keep_end"``: keep the last ``max_packed_len`` tokens.
|
||||
Args:
|
||||
max_seq_len (int): Maximum sequence length. Defaults to 2048.
|
||||
min_chars (int): Minimum number of characters to keep. Defaults to 50.
|
||||
max_chars (int): Maximum number of characters to keep. Defaults to 2_000_000.
|
||||
max_items (Optional[int]): Maximum number of items to process, None=unlimited. Defaults to None.
|
||||
batch_size (int): Number of records tokenized together. Defaults to 256.
|
||||
packing_strategy (str): How to pack sequences: 'simple', 'bfd', or 'bfd_split'. Defaults to "simple".
|
||||
max_packed_len (int): Maximum length of a packed bin. Defaults to 8192.
|
||||
truncation_mode (str): How to truncate over-length sequences: 'keep_start' or 'keep_end'. Defaults to "keep_start".
|
||||
"""
|
||||
|
||||
max_seq_len: int = 2048
|
||||
@@ -72,27 +67,45 @@ class ProcessingConfig(BaseConfig):
|
||||
max_packed_len: int = 8192
|
||||
truncation_mode: str = "keep_start"
|
||||
|
||||
@field_validator("packing_strategy")
|
||||
def _validate_packing_strategy(cls, v: str) -> str:
|
||||
if v not in _PACKING_STRATEGIES:
|
||||
raise ValueError(
|
||||
f"packing_strategy must be one of {sorted(_PACKING_STRATEGIES)}, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
@field_validator("truncation_mode")
|
||||
def _validate_truncation_mode(cls, v: str) -> str:
|
||||
if v not in _TRUNCATION_MODES:
|
||||
raise ValueError(
|
||||
f"truncation_mode must be one of {sorted(_TRUNCATION_MODES)}, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
@field_validator("max_seq_len", "batch_size", "max_packed_len")
|
||||
def _validate_positive_int(cls, v: int) -> int:
|
||||
if v <= 0:
|
||||
raise ValueError(f"must be positive, got {v}")
|
||||
return v
|
||||
|
||||
@field_validator("min_chars")
|
||||
def _validate_non_negative(cls, v: int) -> int:
|
||||
if v < 0:
|
||||
raise ValueError(f"min_chars must be non-negative, got {v}")
|
||||
return v
|
||||
|
||||
|
||||
@dataclass
|
||||
class OutputConfig(BaseConfig):
|
||||
"""Output configuration.
|
||||
"""Output configuration for storage.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
domain_key : Optional[str]
|
||||
Domain key for the output store (default: None).
|
||||
storage_format : str
|
||||
Storage format, one of ``"bin"``, ``"jsonl"`` (default: ``"bin"``).
|
||||
max_tokens_per_shard : int
|
||||
Maximum tokens per shard before splitting (default: 100_000_000).
|
||||
dtype : Dict[str, str]
|
||||
Per-key dtype overrides, e.g. ``{"input_ids": "int32"}`` (default: {}).
|
||||
position_ids_mode : Optional[str]
|
||||
How to compute position_ids in packed sequences.
|
||||
|
||||
- ``"none"``: do not generate (default).
|
||||
- ``"doc_reset"``: reset to 0 at each document boundary.
|
||||
- ``"continuous"``: sequential 0, 1, 2, ... (pretrain, single doc).
|
||||
Args:
|
||||
domain_key (Optional[str]): Domain key for the output store. Defaults to None.
|
||||
storage_format (str): Storage format: 'bin' or 'jsonl'. Defaults to "bin".
|
||||
max_tokens_per_shard (int): Maximum tokens per shard before splitting. Defaults to 100_000_000.
|
||||
dtype (Dict[str, str]): Per-key dtype overrides, e.g. {"input_ids": "int32"}. Defaults to {}.
|
||||
position_ids_mode (str): Position ids mode: 'none', 'doc_reset', or 'continuous'. Defaults to "doc_reset".
|
||||
"""
|
||||
|
||||
domain_key: Optional[str] = None
|
||||
@@ -101,9 +114,36 @@ class OutputConfig(BaseConfig):
|
||||
dtype: Dict[str, str] = field(default_factory=dict)
|
||||
position_ids_mode: str = "doc_reset"
|
||||
|
||||
@field_validator("storage_format")
|
||||
def _validate_storage_format(cls, v: str) -> str:
|
||||
if v not in _STORAGE_FORMATS:
|
||||
raise ValueError(
|
||||
f"storage_format must be one of {sorted(_STORAGE_FORMATS)}, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
@field_validator("position_ids_mode")
|
||||
def _validate_position_ids_mode(cls, v: str) -> str:
|
||||
if v not in _POSITION_IDS_MODES:
|
||||
raise ValueError(
|
||||
f"position_ids_mode must be one of {sorted(_POSITION_IDS_MODES)}, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
|
||||
@dataclass
|
||||
class PipelineConfig(BaseConfig):
|
||||
"""Top-level preprocessing pipeline config.
|
||||
|
||||
Args:
|
||||
version (int): Config schema version. Defaults to 1.
|
||||
input (InputConfig): Input mapping config.
|
||||
mask (Dict[str, str]): Per-field mask labels, e.g. {"system": "mask", "assistant": "train"}. Defaults to {}.
|
||||
mask_default (str): Default mask label for unlisted fields. Defaults to "mask".
|
||||
preprocessing (ProcessingConfig): Processing config.
|
||||
output (OutputConfig): Output config.
|
||||
"""
|
||||
|
||||
version: int = 1
|
||||
input: InputConfig = field(default_factory=InputConfig)
|
||||
mask: Dict[str, str] = field(default_factory=dict)
|
||||
|
||||
+187
-158
@@ -1,7 +1,9 @@
|
||||
from dataclasses import dataclass, field, fields
|
||||
from dataclasses import field
|
||||
from typing import Any, Callable, Dict, List, Optional
|
||||
|
||||
import torch.nn as nn
|
||||
from pydantic import ConfigDict, field_validator, model_validator
|
||||
from pydantic.dataclasses import dataclass
|
||||
from torch.optim import Optimizer
|
||||
from torch.optim.lr_scheduler import LRScheduler
|
||||
from torch.utils.data import Dataset
|
||||
@@ -9,173 +11,200 @@ from torch.utils.data import Dataset
|
||||
from astrai.config.base import BaseConfig
|
||||
from astrai.model.components.lora import LoRAConfig
|
||||
|
||||
|
||||
def required(**kw):
|
||||
return {"required": True, **kw}
|
||||
_TRAIN_TYPES = frozenset({"seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"})
|
||||
_PARALLEL_MODES = frozenset({"none", "ddp", "fsdp"})
|
||||
_BACKENDS = frozenset({"nccl", "gloo"})
|
||||
_START_METHODS = frozenset({"spawn", "fork", "forkserver"})
|
||||
_COMPILE_MODES = frozenset({"default", "reduce-overhead", "max-autotune"})
|
||||
|
||||
|
||||
@dataclass
|
||||
@dataclass(config=ConfigDict(arbitrary_types_allowed=True))
|
||||
class TrainConfig(BaseConfig):
|
||||
# basic setting
|
||||
model_fn: Callable[[], nn.Module] = field(
|
||||
default=None, metadata=required(help="Model factory for training.")
|
||||
)
|
||||
strategy: str = field(default=None, metadata=required(help="Training strategy."))
|
||||
dataset: Dataset = field(
|
||||
default=None, metadata=required(help="Dataset for training.")
|
||||
)
|
||||
optimizer_fn: Callable[[nn.Module], Optimizer] = field(
|
||||
default=None, metadata=required(help="Optimizer factory for training.")
|
||||
)
|
||||
scheduler_fn: Callable[[Optimizer], LRScheduler] = field(
|
||||
default=None, metadata=required(help="Scheduler factory for training.")
|
||||
)
|
||||
n_epoch: int = field(default=1, metadata={"help": "Number of epochs for training."})
|
||||
batch_per_device: int = field(
|
||||
default=4, metadata={"help": "Batch size per device."}
|
||||
)
|
||||
grad_accum_steps: int = field(
|
||||
default=1, metadata={"help": "Number of iterations between steps."}
|
||||
)
|
||||
max_grad_norm: Optional[float] = field(
|
||||
default=1.0,
|
||||
metadata={"help": "Maximum gradient norm. None disables clipping."},
|
||||
)
|
||||
gradient_checkpointing_modules: List[str] = field(
|
||||
default_factory=list,
|
||||
metadata={"help": "Module types to enable activation checkpointing for."},
|
||||
)
|
||||
"""Training configuration.
|
||||
|
||||
# checkpoint setting
|
||||
start_epoch: int = field(default=0, metadata={"help": "Start epoch for training."})
|
||||
start_samples: int = field(
|
||||
default=0,
|
||||
metadata={
|
||||
"help": "Start samples count (per rank). Superseded by checkpoint consumed_samples."
|
||||
},
|
||||
)
|
||||
ckpt_dir: str = field(
|
||||
default="./checkpoint", metadata={"help": "Checkpoint directory."}
|
||||
)
|
||||
ckpt_interval: int = field(
|
||||
default=5000,
|
||||
metadata={"help": "Number of optimizer steps between checkpoints."},
|
||||
)
|
||||
Combines hyperparameters with runtime objects (model_fn, dataset, etc.).
|
||||
Only JSON-serializable fields are written to checkpoint meta via to_dict().
|
||||
|
||||
# lora setting
|
||||
lora: Optional[LoRAConfig] = field(
|
||||
default=None,
|
||||
metadata={"help": "LoRA config. None means full fine-tuning."},
|
||||
)
|
||||
Args:
|
||||
model_fn (Callable[[], nn.Module]): Model factory for training.
|
||||
strategy (str): Training strategy (seq, sft, dpo, grpo, online_*).
|
||||
dataset (Dataset): Dataset for training.
|
||||
optimizer_fn (Callable[[nn.Module], Optimizer]): Optimizer factory for training.
|
||||
scheduler_fn (Callable[[Optimizer], LRScheduler]): Scheduler factory for training.
|
||||
n_epoch (int): Number of epochs for training. Defaults to 1.
|
||||
batch_per_device (int): Batch size per device. Defaults to 4.
|
||||
grad_accum_steps (int): Number of iterations between optimizer steps. Defaults to 1.
|
||||
max_grad_norm (Optional[float]): Maximum gradient norm. None disables clipping. Defaults to 1.0.
|
||||
gradient_checkpointing_modules (List[type]): Module types to enable activation checkpointing for. Defaults to [].
|
||||
compile_mode (Optional[str]): torch.compile mode: 'default', 'reduce-overhead', 'max-autotune', or None. Defaults to None.
|
||||
start_epoch (int): Start epoch for training. Defaults to 0.
|
||||
start_samples (int): Start samples count (per rank). Superseded by checkpoint consumed_samples. Defaults to 0.
|
||||
ckpt_dir (str): Checkpoint directory. Defaults to "./checkpoint".
|
||||
ckpt_interval (int): Number of optimizer steps between checkpoints. Defaults to 5000.
|
||||
lora (Optional[LoRAConfig]): LoRA config. None means full fine-tuning. Defaults to None.
|
||||
metrics (List[str]): Metrics to record during training. Defaults to ["loss", "lr", "grad_norm"].
|
||||
random_seed (int): Random seed. Defaults to 3407.
|
||||
num_workers (int): Number of workers for dataloader. Defaults to 0.
|
||||
prefetch_factor (Optional[int]): Prefetch factor for dataloader. Defaults to None.
|
||||
pin_memory (bool): Pin memory for dataloader. Defaults to False.
|
||||
collate_fn (Optional[Callable[[List[Any]], Any]]): Collate function for dataloader (e.g. dpo_collate_fn). Defaults to None.
|
||||
nprocs (int): Number of processes for distributed training. Defaults to 1.
|
||||
backend (str): Distributed training backend. Defaults to "nccl".
|
||||
master_addr (str): Master address for distributed training. Defaults to "localhost".
|
||||
master_port (str): Master port for distributed training. Defaults to "29500".
|
||||
parallel_mode (str): Parallel strategy: none, ddp, fsdp. Defaults to "none".
|
||||
start_method (str): Multiprocessing start method: spawn/fork/forkserver. Defaults to "spawn".
|
||||
device_type (str): Device type for distributed training. Defaults to "cuda".
|
||||
val_dataset (Optional[Dataset]): Dataset for validation. Defaults to None.
|
||||
val_split (Optional[float]): Ratio to split from training dataset for validation, e.g. 0.05. Defaults to None.
|
||||
val_step (int): Number of optimizer steps between validation runs. Defaults to 1000.
|
||||
neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0.
|
||||
rollout_interval (int): Number of optimizer steps between online rollouts. Defaults to 512.
|
||||
rollout_temperature (float): Sampling temperature for online rollout. Defaults to 0.7.
|
||||
rollout_top_k (int): Top-k filtering for online rollout, 0=disable. Defaults to 0.
|
||||
rollout_top_p (float): Top-p (nucleus) filtering for online rollout. Defaults to 0.9.
|
||||
rollout_max_tokens (int): Maximum generated tokens per response in rollout. Defaults to 1024.
|
||||
reward_model_fn (Optional[Callable]): Factory for reward model, required for online RL strategies. Defaults to None.
|
||||
executor_kwargs (Dict[str, Any]): Extra kwargs passed to ExecutorFactory.create(). Defaults to {}.
|
||||
extra_kwargs (Dict[str, Any]): Other arguments. Defaults to {}.
|
||||
"""
|
||||
|
||||
# metric setting
|
||||
log_dir: str = field(
|
||||
default="./checkpoint/logs", metadata={"help": "Directory for metric logs."}
|
||||
)
|
||||
metrics: List[str] = field(
|
||||
default_factory=lambda: ["loss", "lr", "grad_norm"],
|
||||
metadata={"help": "Metrics to record during training."},
|
||||
)
|
||||
model_fn: Callable[[], nn.Module]
|
||||
strategy: str
|
||||
dataset: Dataset
|
||||
optimizer_fn: Callable[[nn.Module], Optimizer]
|
||||
scheduler_fn: Callable[[Optimizer], LRScheduler]
|
||||
n_epoch: int = 1
|
||||
batch_per_device: int = 4
|
||||
grad_accum_steps: int = 1
|
||||
max_grad_norm: Optional[float] = 1.0
|
||||
gradient_checkpointing_modules: List[type] = field(default_factory=list)
|
||||
compile_mode: Optional[str] = None
|
||||
|
||||
# dataloader setting
|
||||
random_seed: int = field(default=3407, metadata={"help": "Random seed."})
|
||||
num_workers: int = field(
|
||||
default=0, metadata={"help": "Number of workers for dataloader."}
|
||||
)
|
||||
prefetch_factor: Optional[int] = field(
|
||||
default=None, metadata={"help": "Prefetch factor for dataloader."}
|
||||
)
|
||||
pin_memory: bool = field(
|
||||
default=False, metadata={"help": "Pin memory for dataloader."}
|
||||
)
|
||||
collate_fn: Optional[Callable[[List[Any]], Any]] = field(
|
||||
default=None,
|
||||
metadata={"help": "Collate function for dataloader (e.g. dpo_collate_fn)."},
|
||||
)
|
||||
start_epoch: int = 0
|
||||
start_samples: int = 0
|
||||
ckpt_dir: str = "./checkpoint"
|
||||
ckpt_interval: int = 5000
|
||||
|
||||
# distributed training
|
||||
nprocs: int = field(
|
||||
default=1, metadata={"help": "Number of processes for distributed training."}
|
||||
)
|
||||
backend: str = field(
|
||||
default="nccl", metadata={"help": "Distributed training backend."}
|
||||
)
|
||||
master_addr: str = field(
|
||||
default="localhost",
|
||||
metadata={"help": "Master address for distributed training."},
|
||||
)
|
||||
master_port: str = field(
|
||||
default="29500", metadata={"help": "Master port for distributed training."}
|
||||
)
|
||||
parallel_mode: str = field(
|
||||
default="none",
|
||||
metadata={"help": "Parallel strategy: none, ddp, fsdp."},
|
||||
)
|
||||
start_method: str = field(
|
||||
default="spawn",
|
||||
metadata={"help": "Multiprocessing start method (spawn/fork/forkserver)."},
|
||||
)
|
||||
lora: Optional[LoRAConfig] = None
|
||||
|
||||
# others
|
||||
device_type: str = field(
|
||||
default="cuda", metadata={"help": "Device type for distributed training."}
|
||||
)
|
||||
val_dataset: Optional[Dataset] = field(
|
||||
default=None, metadata={"help": "Dataset for validation."}
|
||||
)
|
||||
val_split: Optional[float] = field(
|
||||
default=None,
|
||||
metadata={
|
||||
"help": "Ratio to split from training dataset for validation (e.g. 0.05). Ignored if val_dataset is set."
|
||||
},
|
||||
)
|
||||
val_step: int = field(
|
||||
default=1000,
|
||||
metadata={"help": "Number of optimizer steps between validation runs."},
|
||||
)
|
||||
neftune_alpha: float = field(
|
||||
default=0.0,
|
||||
metadata={"help": "NEFTune noise alpha (0=disabled, typical: 5.0)."},
|
||||
)
|
||||
metrics: List[str] = field(default_factory=lambda: ["loss", "lr", "grad_norm"])
|
||||
|
||||
# online rollout
|
||||
rollout_interval: int = field(
|
||||
default=512,
|
||||
metadata={"help": "Number of optimizer steps between online rollouts."},
|
||||
)
|
||||
rollout_temperature: float = field(
|
||||
default=0.7, metadata={"help": "Sampling temperature for online rollout."}
|
||||
)
|
||||
rollout_top_k: int = field(
|
||||
default=0, metadata={"help": "Top-k filtering for online rollout (0=disable)."}
|
||||
)
|
||||
rollout_top_p: float = field(
|
||||
default=0.9,
|
||||
metadata={"help": "Top-p (nucleus) filtering for online rollout."},
|
||||
)
|
||||
rollout_max_tokens: int = field(
|
||||
default=1024,
|
||||
metadata={"help": "Maximum generated tokens per response in rollout."},
|
||||
)
|
||||
reward_model_fn: Optional[Callable] = field(
|
||||
default=None,
|
||||
metadata={
|
||||
"help": "Factory for reward model (required for online RL strategies)."
|
||||
},
|
||||
)
|
||||
random_seed: int = 3407
|
||||
num_workers: int = 0
|
||||
prefetch_factor: Optional[int] = None
|
||||
pin_memory: bool = False
|
||||
collate_fn: Optional[Callable[[List[Any]], Any]] = None
|
||||
|
||||
executor_kwargs: Dict[str, Any] = field(
|
||||
default_factory=dict,
|
||||
metadata={"help": "Extra kwargs passed to ExecutorFactory.create()."},
|
||||
)
|
||||
extra_kwargs: Dict[str, Any] = field(
|
||||
default_factory=dict, metadata={"help": "Other arguments."}
|
||||
)
|
||||
nprocs: int = 1
|
||||
backend: str = "nccl"
|
||||
master_addr: str = "localhost"
|
||||
master_port: str = "29500"
|
||||
parallel_mode: str = "none"
|
||||
start_method: str = "spawn"
|
||||
|
||||
def __post_init__(self):
|
||||
self.validate()
|
||||
device_type: str = "cuda"
|
||||
val_dataset: Optional[Dataset] = None
|
||||
val_split: Optional[float] = None
|
||||
val_step: int = 1000
|
||||
neftune_alpha: float = 0.0
|
||||
|
||||
def validate(self):
|
||||
for fld in fields(self):
|
||||
if fld.metadata.get("required") and getattr(self, fld.name) is None:
|
||||
raise ValueError(f"TrainConfig.{fld.name} is required but got None.")
|
||||
rollout_interval: int = 512
|
||||
rollout_temperature: float = 0.7
|
||||
rollout_top_k: int = 0
|
||||
rollout_top_p: float = 0.9
|
||||
rollout_max_tokens: int = 1024
|
||||
reward_model_fn: Optional[Callable] = None
|
||||
|
||||
executor_kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||
extra_kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
@field_validator("strategy")
|
||||
def _validate_strategy(cls, v: str) -> str:
|
||||
if v not in _TRAIN_TYPES:
|
||||
raise ValueError(
|
||||
f"strategy must be one of {sorted(_TRAIN_TYPES)}, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
@field_validator("parallel_mode")
|
||||
def _validate_parallel_mode(cls, v: str) -> str:
|
||||
if v not in _PARALLEL_MODES:
|
||||
raise ValueError(
|
||||
f"parallel_mode must be one of {sorted(_PARALLEL_MODES)}, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
@field_validator("backend")
|
||||
def _validate_backend(cls, v: str) -> str:
|
||||
if v not in _BACKENDS:
|
||||
raise ValueError(f"backend must be one of {sorted(_BACKENDS)}, got {v!r}")
|
||||
return v
|
||||
|
||||
@field_validator("start_method")
|
||||
def _validate_start_method(cls, v: str) -> str:
|
||||
if v not in _START_METHODS:
|
||||
raise ValueError(
|
||||
f"start_method must be one of {sorted(_START_METHODS)}, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
@field_validator("compile_mode")
|
||||
def _validate_compile_mode(cls, v: Optional[str]) -> Optional[str]:
|
||||
if v is not None and v not in _COMPILE_MODES:
|
||||
raise ValueError(
|
||||
f"compile_mode must be one of {sorted(_COMPILE_MODES)} or None, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
@field_validator(
|
||||
"n_epoch",
|
||||
"batch_per_device",
|
||||
"grad_accum_steps",
|
||||
"ckpt_interval",
|
||||
"val_step",
|
||||
"rollout_interval",
|
||||
"rollout_max_tokens",
|
||||
)
|
||||
def _validate_positive_int(cls, v: int) -> int:
|
||||
if v <= 0:
|
||||
raise ValueError(f"must be positive, got {v}")
|
||||
return v
|
||||
|
||||
@field_validator("rollout_temperature")
|
||||
def _validate_positive_float(cls, v: float) -> float:
|
||||
if v <= 0:
|
||||
raise ValueError(f"must be positive, got {v}")
|
||||
return v
|
||||
|
||||
@field_validator("rollout_top_p")
|
||||
def _validate_top_p(cls, v: float) -> float:
|
||||
if not 0 < v <= 1:
|
||||
raise ValueError(f"rollout_top_p must be in (0, 1], got {v}")
|
||||
return v
|
||||
|
||||
@field_validator("rollout_top_k", "num_workers", "neftune_alpha")
|
||||
def _validate_non_negative(cls, v):
|
||||
if v < 0:
|
||||
raise ValueError(f"must be non-negative, got {v}")
|
||||
return v
|
||||
|
||||
@field_validator("max_grad_norm")
|
||||
def _validate_max_grad_norm(cls, v: Optional[float]) -> Optional[float]:
|
||||
if v is not None and v <= 0:
|
||||
raise ValueError(f"max_grad_norm must be positive or None, got {v}")
|
||||
return v
|
||||
|
||||
@field_validator("val_split")
|
||||
def _validate_val_split(cls, v: Optional[float]) -> Optional[float]:
|
||||
if v is not None and not 0 < v < 1:
|
||||
raise ValueError(f"val_split must be in (0, 1) or None, got {v}")
|
||||
return v
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _validate_online_strategy(self) -> "TrainConfig":
|
||||
if self.strategy.startswith("online_") and self.reward_model_fn is None:
|
||||
raise ValueError(
|
||||
f"reward_model_fn is required for online RL strategy {self.strategy!r}"
|
||||
)
|
||||
return self
|
||||
|
||||
@@ -6,7 +6,6 @@ from astrai.dataset.dataset import (
|
||||
)
|
||||
from astrai.dataset.sampler import RDSampler
|
||||
from astrai.dataset.storage import (
|
||||
H5Store,
|
||||
JsonlStore,
|
||||
MmapStore,
|
||||
Recordable,
|
||||
@@ -17,9 +16,7 @@ from astrai.dataset.storage import (
|
||||
)
|
||||
from astrai.serialization import (
|
||||
load_bin,
|
||||
load_h5,
|
||||
save_bin,
|
||||
save_h5,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
@@ -31,12 +28,9 @@ __all__ = [
|
||||
"Streamable",
|
||||
"Recordable",
|
||||
"StoreFactory",
|
||||
"H5Store",
|
||||
"MmapStore",
|
||||
"JsonlStore",
|
||||
"detect_format",
|
||||
"save_h5",
|
||||
"load_h5",
|
||||
"save_bin",
|
||||
"load_bin",
|
||||
"RDSampler",
|
||||
|
||||
@@ -314,7 +314,7 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
|
||||
stream datasets (SEQ/SFT). Record datasets ignore it.
|
||||
stride: Stride between consecutive stream samples
|
||||
(default: same as *window_size*).
|
||||
storage_type: Storage backend ("h5", "bin", "jsonl") or
|
||||
storage_type: Storage backend ("bin", "jsonl") or
|
||||
None for auto-detection.
|
||||
tokenizer_path: Path to tokenizer for lazy JSONL
|
||||
tokenisation (record datasets only).
|
||||
@@ -384,7 +384,7 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
|
||||
"""Build an on-the-fly tokenisation processor if applicable.
|
||||
|
||||
Only raw JSONL + record datasets (DPO/GRPO) need a processor;
|
||||
pre-tokenised backends (H5/bin) and stream datasets (SEQ/SFT)
|
||||
pre-tokenised backends (bin) and stream datasets (SEQ/SFT)
|
||||
return ``None`` so no tokenizer is loaded.
|
||||
"""
|
||||
if tokenizer_path is None or storage_type != "jsonl":
|
||||
@@ -451,7 +451,7 @@ class DPODataset(BaseDataset):
|
||||
|
||||
Two loading paths (handled by :class:`DatasetFactory`):
|
||||
|
||||
- **Pre-tokenized** (H5/bin): ``store.load(path)`` reads per-record
|
||||
- **Pre-tokenized** (bin): ``store.load(path)`` reads per-record
|
||||
tensors; ``__getitem__`` returns them directly.
|
||||
- **Raw JSONL** (``tokenizer_path=...``): builds a lazy processor
|
||||
via :func:`dpo_processor` that tokenises on the fly — no packing,
|
||||
|
||||
@@ -10,7 +10,6 @@ Architecture (composition over inheritance):
|
||||
Streamable (mixin) — raw token slice fetch(begin, end, keys)
|
||||
Recordable (mixin) — raw record slice fetch_record(idx, keys)
|
||||
|
||||
H5Store(Store, Streamable, Recordable)
|
||||
MmapStore(Store, Streamable, Recordable)
|
||||
JsonlStore(Store, Streamable, Recordable)
|
||||
|
||||
@@ -36,9 +35,9 @@ control. ``store.token_count`` is the total stream token count (what
|
||||
``len(store)`` used to mean in the legacy stream-only API).
|
||||
|
||||
``segments_are_records`` (class attribute on each Store subclass)
|
||||
tells ``_normalize`` whether segments are inherently per-record (H5/
|
||||
JSONL) or opaque shards (bin). Record access for bin relies on
|
||||
``_offsets`` instead.
|
||||
tells ``_normalize`` whether segments are inherently per-record (JSONL)
|
||||
or opaque shards (bin). Record access for bin relies on ``_offsets``
|
||||
instead.
|
||||
|
||||
:class:`JsonlStore` supports a lazy mode (``processor=fn``) that keeps
|
||||
raw records and defers tokenisation to ``fetch_record`` — used by DPO
|
||||
@@ -62,7 +61,6 @@ from astrai.preprocessing.transform import TokenizeTransform
|
||||
from astrai.serialization import (
|
||||
load_bin,
|
||||
load_bin_offsets,
|
||||
load_h5,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -83,19 +81,10 @@ def detect_format(load_path: str) -> str:
|
||||
root = Path(load_path)
|
||||
if root.is_file():
|
||||
suffix = root.suffix.lower()
|
||||
if suffix in (".h5", ".hdf5"):
|
||||
return "h5"
|
||||
if suffix == ".jsonl":
|
||||
return "jsonl"
|
||||
raise ValueError(f"Unsupported file format: {suffix}")
|
||||
|
||||
h5_files = [
|
||||
Path(p)
|
||||
for pattern in ("*.h5", "*.hdf5")
|
||||
for p in glob.glob(str(root / "**" / pattern), recursive=True)
|
||||
]
|
||||
if h5_files:
|
||||
return "h5"
|
||||
bin_files = [Path(p) for p in glob.glob(str(root / "**" / "*.bin"), recursive=True)]
|
||||
if bin_files:
|
||||
has_meta = (root / "meta.json").exists() or len(
|
||||
@@ -185,7 +174,7 @@ class Store(ABC):
|
||||
"""Number of records available via :meth:`fetch_record`.
|
||||
|
||||
Non-zero only when the backing layout provides per-record
|
||||
indexing (H5/JSONL segments or bin ``_offsets``).
|
||||
indexing (JSONL segments or bin ``_offsets``).
|
||||
"""
|
||||
return self._num_records
|
||||
|
||||
@@ -269,7 +258,7 @@ class Store(ABC):
|
||||
Record mode: if *offsets* is provided (bin layout),
|
||||
``_offsets[key]`` stores cumulative per-record offsets into the
|
||||
single concatenated segment. Otherwise, when
|
||||
``segments_are_records`` is True (H5/JSONL), ``_data[key]`` is
|
||||
``segments_are_records`` is True (JSONL), ``_data[key]`` is
|
||||
a per-record list and ``fetch_record`` indexes it directly.
|
||||
|
||||
Nested keys (GRPO ``responses``/``masks`` as
|
||||
@@ -305,7 +294,7 @@ class Store(ABC):
|
||||
logger.warning(
|
||||
"Key '%s' has %d segments with offsets — record mode "
|
||||
"disabled for this key (multi-shard bin+offsets not "
|
||||
"supported). Merge shards or use H5/JSONL.",
|
||||
"supported). Merge shards or use JSONL.",
|
||||
key,
|
||||
len(segs),
|
||||
)
|
||||
@@ -330,7 +319,7 @@ class Streamable:
|
||||
Stateless trait relying on ``self._data``, ``self._cum``,
|
||||
``self._length`` maintained by :class:`Store`. Stream mode is
|
||||
active when the owning store has ``window_size > 0``; for stores
|
||||
that can also serve record access (H5/JSONL/bin+offsets), the
|
||||
that can also serve record access (JSONL/bin+offsets), the
|
||||
``fetch_record`` API from :class:`Recordable` is used instead.
|
||||
"""
|
||||
|
||||
@@ -415,33 +404,6 @@ class StoreFactory(BaseFactory["Store"]):
|
||||
"""Factory for creating Store instances by type name."""
|
||||
|
||||
|
||||
@StoreFactory.register("h5")
|
||||
class H5Store(Store, Streamable, Recordable):
|
||||
"""HDF5-based storage backend (pre-tokenized data).
|
||||
|
||||
Each key is stored as a group of per-record datasets (``data_0``,
|
||||
``data_1``, …). Supports both access modes:
|
||||
|
||||
- **Stream**: ``fetch(begin, end, key)`` and ``store[i]`` slice
|
||||
across concatenated records via ``_cum`` — used by SEQ/SFT.
|
||||
- **Record**: ``fetch_record(i, key)`` and ``store[i]`` (when
|
||||
``window_size == 0``) index ``_data[key]`` directly — used by
|
||||
DPO/GRPO.
|
||||
"""
|
||||
|
||||
segments_are_records = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
window_size: int = 0,
|
||||
stride: Optional[int] = None,
|
||||
):
|
||||
super().__init__(window_size=window_size, stride=stride)
|
||||
|
||||
def load(self, path: str, **kwargs):
|
||||
self._normalize(load_h5(path))
|
||||
|
||||
|
||||
@StoreFactory.register("bin")
|
||||
class MmapStore(Store, Streamable, Recordable):
|
||||
"""Memory-mapped binary storage backend.
|
||||
|
||||
@@ -4,27 +4,44 @@ Public API:
|
||||
- ``attn_decode`` — single-query decode attention
|
||||
- ``attn_prefill`` — multi-query prefill attention
|
||||
- ``attn_paged_decode`` — paged decode attention (direct page-table access)
|
||||
- ``AttentionBackend`` — ABC for attention computation strategies
|
||||
- ``TorchNativeBackend`` — default SDPA backend with KV cache I/O
|
||||
- ``CudaBackend`` — CUDA kernel backend with paged decode + prefill
|
||||
|
||||
Interface (shared by all wrappers):
|
||||
causal_offset: -1 = non-causal; >=0 = absolute position of first Q token
|
||||
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True = keep)
|
||||
scale: 0.0 = auto (1/sqrt(head_dim)); >0 = explicit
|
||||
layout: "bhld" (default) or "blhd"
|
||||
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
||||
(blhd). Scale is always ``1/sqrt(head_dim)``.
|
||||
|
||||
Causal and mask can coexist — both are applied simultaneously.
|
||||
|
||||
Each wrapper dispatches to its compiled CUDA kernel (``astrai.extension.attn_*``)
|
||||
when available, otherwise falls back to ``torch.nn.functional.scaled_dot_product_attention``.
|
||||
Each wrapper calls its compiled CUDA kernel directly. Fallback to torch
|
||||
SDPA is handled by the attention backend, not the wrapper functions.
|
||||
"""
|
||||
|
||||
from astrai.extension.attention_backend import (
|
||||
ATTN_BACKEND,
|
||||
AttentionBackend,
|
||||
CudaBackend,
|
||||
TorchNativeBackend,
|
||||
attention,
|
||||
attn_backend,
|
||||
get_backend,
|
||||
)
|
||||
from astrai.extension.attention_ops import (
|
||||
attn_decode,
|
||||
attn_paged_decode,
|
||||
attn_prefill,
|
||||
)
|
||||
from astrai.extension.loader import KERNEL_NAMES, is_available
|
||||
from astrai.extension.ops import attention, attn_decode, attn_paged_decode, attn_prefill
|
||||
|
||||
__all__ = [
|
||||
"ATTN_BACKEND",
|
||||
"AttentionBackend",
|
||||
"CudaBackend",
|
||||
"TorchNativeBackend",
|
||||
"attention",
|
||||
"attn_backend",
|
||||
"get_backend",
|
||||
"attn_decode",
|
||||
"attn_paged_decode",
|
||||
"attn_prefill",
|
||||
"attention",
|
||||
"is_available",
|
||||
"KERNEL_NAMES",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,410 @@
|
||||
"""Attention backend abstraction with context-manager switching.
|
||||
|
||||
The backend encapsulates KV cache I/O and attention computation. The
|
||||
attention module (GQA/MLA) keeps projections, rotary, QK-norm, gating,
|
||||
and output projection; the backend handles everything from "write K/V
|
||||
to cache" through "SDPA output".
|
||||
|
||||
Usage — mirroring ``torch.nn.attention.sdpa_kernel``:
|
||||
|
||||
from astrai.extension import attn_backend, ATTN_BACKEND
|
||||
|
||||
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
|
||||
engine.generate("hello")
|
||||
|
||||
# or with an instance:
|
||||
with attn_backend(TorchNativeBackend()):
|
||||
...
|
||||
|
||||
# or the shorthand (instance is itself a context manager):
|
||||
with TorchNativeBackend():
|
||||
...
|
||||
|
||||
Thread-safe via ``contextvars`` — each scheduler thread gets its own
|
||||
active backend. ``get_backend()`` returns the active one, falling back
|
||||
to a process-wide ``TorchNativeBackend`` singleton.
|
||||
|
||||
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
||||
(blhd). The backend returns ``[batch, seq_len, n_heads * head_dim]``.
|
||||
"""
|
||||
|
||||
import contextvars
|
||||
import enum
|
||||
import math
|
||||
from abc import ABC, abstractmethod
|
||||
from contextlib import contextmanager
|
||||
from typing import Optional, Union
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.extension.attention_ops import attn_paged_decode, attn_prefill
|
||||
from astrai.extension.loader import is_available
|
||||
from astrai.inference.core.cache import KVCache
|
||||
|
||||
_current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.ContextVar(
|
||||
"attn_backend"
|
||||
)
|
||||
|
||||
|
||||
class ATTN_BACKEND(enum.Enum):
|
||||
"""Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``."""
|
||||
|
||||
TORCH_NATIVE = "torch_native"
|
||||
CUDA = "cuda"
|
||||
|
||||
|
||||
def get_backend() -> "AttentionBackend":
|
||||
"""Return the active backend for the current thread/context.
|
||||
|
||||
Falls back to a ``TorchNativeBackend`` singleton when no backend
|
||||
has been activated via ``with``.
|
||||
"""
|
||||
try:
|
||||
return _current_backend.get()
|
||||
except LookupError:
|
||||
return _default_backend
|
||||
|
||||
|
||||
@contextmanager
|
||||
def attn_backend(backend: Union[ATTN_BACKEND, "AttentionBackend", type]):
|
||||
"""Context manager to select an attention backend.
|
||||
|
||||
Mirrors ``torch.nn.attention.sdpa_kernel``. Accepts an
|
||||
``ATTN_BACKEND`` enum value, a backend class, or a backend instance.
|
||||
|
||||
Examples::
|
||||
|
||||
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
|
||||
...
|
||||
with attn_backend(TorchNativeBackend):
|
||||
...
|
||||
with attn_backend(TorchNativeBackend()):
|
||||
...
|
||||
"""
|
||||
if isinstance(backend, ATTN_BACKEND):
|
||||
instance = _BACKEND_REGISTRY[backend]()
|
||||
elif isinstance(backend, type) and issubclass(backend, AttentionBackend):
|
||||
instance = backend()
|
||||
elif isinstance(backend, AttentionBackend):
|
||||
instance = backend
|
||||
else:
|
||||
raise TypeError(
|
||||
f"expected ATTN_BACKEND, AttentionBackend type, or instance, "
|
||||
f"got {type(backend).__name__}"
|
||||
)
|
||||
token = _current_backend.set(instance)
|
||||
try:
|
||||
yield instance
|
||||
finally:
|
||||
_current_backend.reset(token)
|
||||
|
||||
|
||||
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
|
||||
"""Expand KV heads to match Q heads for GQA."""
|
||||
bs, slen, n_heads, head_dim = x.shape
|
||||
if n_rep == 1:
|
||||
return x
|
||||
return (
|
||||
x[:, :, :, None, :]
|
||||
.expand(bs, slen, n_heads, n_rep, head_dim)
|
||||
.reshape(bs, slen, n_heads * n_rep, head_dim)
|
||||
)
|
||||
|
||||
|
||||
def attention(
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional[KVCache] = None,
|
||||
layer_id: int = 0,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
"""Functional attention entry point — mirrors ``F.scaled_dot_product_attention``.
|
||||
|
||||
Delegates to the active backend (set via ``with attn_backend(...)``).
|
||||
Handles KV cache I/O, GQA head expansion, and causal masking so the
|
||||
caller only needs to provide projected q/k/v.
|
||||
|
||||
Args:
|
||||
q: [batch, q_len, n_heads, head_dim] (blhd)
|
||||
k: [batch, q_len, n_kv_heads, head_dim] (blhd)
|
||||
v: [batch, q_len, n_kv_heads, head_dim] (blhd)
|
||||
kv_cache: cache dataclass, or None for training (no cache).
|
||||
layer_id: transformer layer index for buffer access.
|
||||
attn_mask: pre-built attention mask (SDPA-compatible).
|
||||
is_causal: whether to apply causal masking.
|
||||
|
||||
Returns:
|
||||
[batch, q_len, n_heads * head_dim]
|
||||
"""
|
||||
backend = get_backend()
|
||||
return backend.forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||
|
||||
|
||||
class AttentionBackend(ABC):
|
||||
"""Abstract base for attention computation strategies.
|
||||
|
||||
Subclasses implement ``fwd_decode`` (q_len == 1, with cache) and
|
||||
``fwd_prefill`` (q_len > 1, with or without cache). The public
|
||||
``forward`` method dispatches based on q_len.
|
||||
|
||||
Three equivalent ways to activate a backend::
|
||||
|
||||
with attn_backend(ATTN_BACKEND.TORCH_NATIVE): # enum
|
||||
...
|
||||
with attn_backend(TorchNativeBackend): # class
|
||||
...
|
||||
with TorchNativeBackend(): # instance
|
||||
...
|
||||
"""
|
||||
|
||||
def __enter__(self) -> "AttentionBackend":
|
||||
self._token = _current_backend.set(self)
|
||||
return self
|
||||
|
||||
def __exit__(self, *exc) -> None:
|
||||
_current_backend.reset(self._token)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional[KVCache],
|
||||
layer_id: int,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
"""Dispatch to decode or extend based on q_len.
|
||||
|
||||
Args:
|
||||
q: [batch, q_len, n_heads, head_dim]
|
||||
k: [batch, q_len, n_kv_heads, head_dim]
|
||||
v: [batch, q_len, n_kv_heads, head_dim]
|
||||
kv_cache: cache dataclass, or None for training (no cache).
|
||||
layer_id: transformer layer index for buffer access.
|
||||
attn_mask: pre-built attention mask compatible with SDPA.
|
||||
is_causal: whether to apply causal masking.
|
||||
|
||||
Returns:
|
||||
[batch, q_len, n_heads * head_dim]
|
||||
"""
|
||||
if kv_cache is not None and q.size(1) == 1:
|
||||
return self.fwd_decode(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||
return self.fwd_prefill(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||
|
||||
@abstractmethod
|
||||
def fwd_decode(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional[KVCache],
|
||||
layer_id: int,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
"""Single-token decode with KV cache."""
|
||||
|
||||
@abstractmethod
|
||||
def fwd_prefill(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional[KVCache],
|
||||
layer_id: int,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
"""Multi-token prefill or training forward."""
|
||||
|
||||
|
||||
class TorchNativeBackend(AttentionBackend):
|
||||
"""Reference backend using torch SDPA with indirect KV cache indexing.
|
||||
|
||||
Writes new K/V into the cache buffers, gathers the full sequence K/V
|
||||
via ``req_to_token`` indirect indexing, then calls
|
||||
``F.scaled_dot_product_attention``.
|
||||
|
||||
For training (``kv_cache is None``), skips cache I/O entirely and
|
||||
runs SDPA directly on the projected q/k/v.
|
||||
"""
|
||||
|
||||
def fwd_decode(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional[KVCache],
|
||||
layer_id: int,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||
|
||||
def fwd_prefill(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional[KVCache],
|
||||
layer_id: int,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||
|
||||
def _forward(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional[KVCache],
|
||||
layer_id: int,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
if kv_cache is not None:
|
||||
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
||||
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||
|
||||
max_len = kv_cache.seq_lens.max()
|
||||
indices = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
||||
pos_mask = (
|
||||
torch.arange(max_len, device=q.device)[None, :]
|
||||
< kv_cache.seq_lens[:, None]
|
||||
)
|
||||
indices = torch.where(pos_mask, indices, torch.zeros_like(indices))
|
||||
k = kv_cache.k_buffer[layer_id, indices]
|
||||
v = kv_cache.v_buffer[layer_id, indices]
|
||||
|
||||
n_rep = q.size(2) // k.size(2)
|
||||
if n_rep > 1:
|
||||
k = repeat_kv(k, n_rep)
|
||||
v = repeat_kv(v, n_rep)
|
||||
|
||||
q = q.permute(0, 2, 1, 3)
|
||||
k = k.permute(0, 2, 1, 3)
|
||||
v = v.permute(0, 2, 1, 3)
|
||||
|
||||
out = F.scaled_dot_product_attention(q, k, v, attn_mask, is_causal=is_causal)
|
||||
out = out.permute(0, 2, 1, 3).contiguous().flatten(2)
|
||||
return out
|
||||
|
||||
|
||||
_default_backend = TorchNativeBackend()
|
||||
|
||||
|
||||
class CudaBackend(AttentionBackend):
|
||||
"""CUDA kernel backend with direct KV cache access.
|
||||
|
||||
Decode path: writes K/V to cache, then calls ``attn_paged_decode``
|
||||
with ``page_size=1`` (each token slot is a single-token "page").
|
||||
The ``req_to_token`` table serves directly as the page table.
|
||||
|
||||
Prefill path: writes K/V to cache, gathers full-sequence K/V via
|
||||
indirect indexing (same as TorchNativeBackend), then calls
|
||||
``attn_prefill``.
|
||||
|
||||
Training path (``kv_cache is None``): calls ``attn_prefill`` directly
|
||||
on the projected q/k/v.
|
||||
|
||||
Falls back to ``TorchNativeBackend`` for any path where the
|
||||
corresponding CUDA kernel is not available.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self._fallback = TorchNativeBackend()
|
||||
|
||||
def fwd_decode(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional[KVCache],
|
||||
layer_id: int,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
if kv_cache is None or not is_available("attn_paged_decode"):
|
||||
return self._fallback.fwd_decode(
|
||||
q, k, v, kv_cache, layer_id, attn_mask, is_causal
|
||||
)
|
||||
|
||||
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
||||
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||
|
||||
seq_lens = kv_cache.seq_lens
|
||||
max_len = kv_cache.max_len
|
||||
|
||||
page_table = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
||||
|
||||
k_cache = kv_cache.k_buffer[layer_id].unsqueeze(1)
|
||||
v_cache = kv_cache.v_buffer[layer_id].unsqueeze(1)
|
||||
|
||||
if q.size(0) == 1:
|
||||
mask = None
|
||||
else:
|
||||
mask = torch.arange(max_len, device=q.device)[None, :] < seq_lens[:, None]
|
||||
|
||||
out = attn_paged_decode(
|
||||
q,
|
||||
page_table,
|
||||
k_cache,
|
||||
v_cache,
|
||||
page_size=1,
|
||||
kv_len=max_len,
|
||||
mask=mask,
|
||||
is_causal=is_causal,
|
||||
)
|
||||
|
||||
out = out.flatten(2)
|
||||
return out
|
||||
|
||||
def fwd_prefill(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional[KVCache],
|
||||
layer_id: int,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
if kv_cache is None:
|
||||
if is_available("attn_prefill"):
|
||||
out = attn_prefill(q, k, v, mask=attn_mask, is_causal=is_causal)
|
||||
return out.flatten(2)
|
||||
return self._fallback.fwd_prefill(
|
||||
q, k, v, kv_cache, layer_id, attn_mask, is_causal
|
||||
)
|
||||
|
||||
if not is_available("attn_prefill"):
|
||||
return self._fallback.fwd_prefill(
|
||||
q, k, v, kv_cache, layer_id, attn_mask, is_causal
|
||||
)
|
||||
|
||||
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
||||
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||
|
||||
max_len = kv_cache.max_len
|
||||
indices = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
||||
pos_mask = (
|
||||
torch.arange(max_len, device=q.device)[None, :] < kv_cache.seq_lens[:, None]
|
||||
)
|
||||
indices = torch.where(pos_mask, indices, torch.zeros_like(indices))
|
||||
k_full = kv_cache.k_buffer[layer_id, indices]
|
||||
v_full = kv_cache.v_buffer[layer_id, indices]
|
||||
|
||||
out = attn_prefill(q, k_full, v_full, mask=attn_mask, is_causal=is_causal)
|
||||
return out.flatten(2)
|
||||
|
||||
|
||||
_BACKEND_REGISTRY: dict[ATTN_BACKEND, type[AttentionBackend]] = {
|
||||
ATTN_BACKEND.TORCH_NATIVE: TorchNativeBackend,
|
||||
ATTN_BACKEND.CUDA: CudaBackend,
|
||||
}
|
||||
@@ -0,0 +1,117 @@
|
||||
"""Attention kernel wrapper functions — one entry point per compiled kernel.
|
||||
|
||||
Each wrapper calls its CUDA kernel directly. If the kernel is not
|
||||
available, raises ``RuntimeError``. Fallback to torch SDPA is the
|
||||
responsibility of the attention backend, not this module.
|
||||
|
||||
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
||||
(blhd). Scale is always ``1/sqrt(head_dim)``.
|
||||
|
||||
Interface (all functions):
|
||||
is_causal: True = causal mask; False = non-causal
|
||||
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.extension.loader import _available, _modules
|
||||
|
||||
|
||||
def _check_available(name: str):
|
||||
if not _available.get(name):
|
||||
raise RuntimeError(
|
||||
f"CUDA kernel '{name}' is not available. "
|
||||
f"Build with CSRC_KERNELS=true or use a torch-native backend."
|
||||
)
|
||||
|
||||
|
||||
def attn_decode(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
mask: torch.Tensor | None = None,
|
||||
is_causal: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""GQA decode attention (q_len == 1).
|
||||
|
||||
Args:
|
||||
q: [batch, 1, n_heads, head_dim] (blhd, bf16)
|
||||
k: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||
v: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||
mask: 2D [batch, kv_len] or 3D [batch, 1, kv_len] (bool, True=keep)
|
||||
is_causal: apply causal mask
|
||||
|
||||
Returns:
|
||||
[batch, 1, n_heads, head_dim] (blhd, bf16)
|
||||
"""
|
||||
_check_available("attn_decode")
|
||||
causal_offset = (k.size(1) - 1) if is_causal else -1
|
||||
return _modules["attn_decode"].attn_decode(
|
||||
q, k, v, mask=mask, causal_offset=causal_offset, layout=1
|
||||
)
|
||||
|
||||
|
||||
def attn_prefill(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
mask: torch.Tensor | None = None,
|
||||
is_causal: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""GQA prefill attention (q_len > 1).
|
||||
|
||||
Args:
|
||||
q: [batch, q_len, n_heads, head_dim] (blhd, bf16)
|
||||
k: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||
v: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
|
||||
is_causal: apply causal mask
|
||||
|
||||
Returns:
|
||||
[batch, q_len, n_heads, head_dim] (blhd, bf16)
|
||||
"""
|
||||
_check_available("attn_prefill")
|
||||
causal_offset = (k.size(1) - q.size(1)) if is_causal else -1
|
||||
return _modules["attn_prefill"].attn_prefill(
|
||||
q, k, v, mask=mask, causal_offset=causal_offset, layout=1
|
||||
)
|
||||
|
||||
|
||||
def attn_paged_decode(
|
||||
q: torch.Tensor,
|
||||
page_table: torch.Tensor,
|
||||
k_cache: torch.Tensor,
|
||||
v_cache: torch.Tensor,
|
||||
page_size: int,
|
||||
kv_len: int,
|
||||
mask: torch.Tensor | None = None,
|
||||
is_causal: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""Paged GQA decode attention (q_len == 1, direct page-table access).
|
||||
|
||||
Args:
|
||||
q: [batch, 1, n_heads, head_dim] (blhd, bf16)
|
||||
page_table: [batch, max_pages] (int64)
|
||||
k_cache: [n_pages, page_size, n_kv_heads, head_dim] (bf16)
|
||||
v_cache: same as k_cache
|
||||
page_size: tokens per page
|
||||
kv_len: actual sequence length per request
|
||||
mask: 2D [batch, kv_len] or 3D [batch, 1, kv_len] (bool, True=keep)
|
||||
is_causal: apply causal mask
|
||||
|
||||
Returns:
|
||||
[batch, 1, n_heads, head_dim] (blhd, bf16)
|
||||
"""
|
||||
_check_available("attn_paged_decode")
|
||||
causal_offset = (kv_len - 1) if is_causal else -1
|
||||
return _modules["attn_paged_decode"].attn_paged_decode(
|
||||
q,
|
||||
page_table,
|
||||
k_cache,
|
||||
v_cache,
|
||||
page_size,
|
||||
kv_len,
|
||||
mask=mask,
|
||||
causal_offset=causal_offset,
|
||||
layout=1,
|
||||
)
|
||||
@@ -1,298 +0,0 @@
|
||||
"""GQA attention wrapper functions — one entry point per compiled kernel.
|
||||
|
||||
Each wrapper dispatches to its CUDA kernel (loaded in ``loader.py``) when
|
||||
available, otherwise falls back to ``torch`` SDPA.
|
||||
|
||||
Interface (all functions):
|
||||
causal_offset: -1 = non-causal; >=0 = absolute position of first Q token
|
||||
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool)
|
||||
scale: 0.0 = auto (1/sqrt(head_dim)); >0 = explicit
|
||||
layout: "bhld" (default) or "blhd"
|
||||
|
||||
Add new kernel wrappers here; split into per-variant files only if this file
|
||||
grows large.
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from astrai.extension.loader import _available, _modules
|
||||
|
||||
_LAYOUT_CODES: dict[str, int] = {"bhld": 0, "blhd": 1}
|
||||
|
||||
|
||||
def _parse_layout(layout: str | int) -> int:
|
||||
if isinstance(layout, int):
|
||||
return layout
|
||||
code = _LAYOUT_CODES.get(layout.lower())
|
||||
if code is None:
|
||||
raise ValueError(
|
||||
f"unknown layout '{layout}', expected one of {list(_LAYOUT_CODES)}"
|
||||
)
|
||||
return code
|
||||
|
||||
|
||||
def _to_bhld(t: torch.Tensor, layout: int) -> torch.Tensor:
|
||||
"""Normalize to b h l d view. Zero-copy transpose if layout==1 (b l h d)."""
|
||||
if layout == 1:
|
||||
return t.transpose(1, 2)
|
||||
return t
|
||||
|
||||
|
||||
def _expand_kv_heads(
|
||||
k: torch.Tensor, v: torch.Tensor, q_head: int
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Expand K/V heads to match Q heads for GQA fallback."""
|
||||
kv_head = k.size(1)
|
||||
if kv_head == q_head:
|
||||
return k, v
|
||||
group = q_head // kv_head
|
||||
k = k.repeat_interleave(group, dim=1)
|
||||
v = v.repeat_interleave(group, dim=1)
|
||||
return k, v
|
||||
|
||||
|
||||
def _build_attn_mask(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
mask: torch.Tensor | None,
|
||||
causal_offset: int,
|
||||
scale: float,
|
||||
) -> tuple[torch.Tensor | None, float]:
|
||||
"""Build SDPA-compatible attn_mask + resolved scale.
|
||||
|
||||
q and k must already be in b h l d layout.
|
||||
Causal and mask can coexist: causal sets -inf above the diagonal, mask
|
||||
sets -inf for padded positions. Both are OR'd into a single bool mask.
|
||||
"""
|
||||
q_len = q.size(2)
|
||||
kv_len = k.size(2)
|
||||
head_dim = q.size(3)
|
||||
resolved_scale = scale if scale and scale > 0 else 1.0 / math.sqrt(head_dim)
|
||||
|
||||
attn_mask = None
|
||||
|
||||
if mask is not None:
|
||||
if mask.dim() == 2:
|
||||
# [batch, kv_len] → [batch, 1, 1, kv_len]
|
||||
attn_mask = mask[:, None, None, :]
|
||||
elif mask.dim() == 3:
|
||||
# [batch, q_len, kv_len] → [batch, 1, q_len, kv_len]
|
||||
attn_mask = mask[:, None, :, :]
|
||||
else:
|
||||
raise ValueError(f"mask must be 2D or 3D, got {mask.dim()}D")
|
||||
|
||||
if causal_offset >= 0:
|
||||
batch = q.size(0)
|
||||
# q row i attends to kv cols 0..(causal_offset + i)
|
||||
q_idx = torch.arange(q_len, device=q.device).unsqueeze(1) # [q_len, 1]
|
||||
kv_idx = torch.arange(kv_len, device=q.device).unsqueeze(0) # [1, kv_len]
|
||||
causal_bool = kv_idx > (causal_offset + q_idx) # True = masked out
|
||||
causal_mask = causal_bool.unsqueeze(0).expand(
|
||||
batch, -1, -1
|
||||
) # [batch, q_len, kv_len]
|
||||
causal_mask = causal_mask[:, None, :, :] # [batch, 1, q_len, kv_len]
|
||||
|
||||
if attn_mask is not None:
|
||||
attn_mask = attn_mask | causal_mask
|
||||
else:
|
||||
attn_mask = causal_mask
|
||||
|
||||
return attn_mask, resolved_scale
|
||||
|
||||
|
||||
def _torch_fallback(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
mask: torch.Tensor | None,
|
||||
causal_offset: int,
|
||||
scale: float,
|
||||
q_layout: int,
|
||||
kv_layout: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Reference attention via ``scaled_dot_product_attention``.
|
||||
|
||||
q_layout / kv_layout: 0 = b h l d, 1 = b l h d.
|
||||
If kv_layout is None, uses q_layout (Q and K/V share the same layout).
|
||||
"""
|
||||
if kv_layout is None:
|
||||
kv_layout = q_layout
|
||||
q = _to_bhld(q, q_layout)
|
||||
k = _to_bhld(k, kv_layout)
|
||||
v = _to_bhld(v, kv_layout)
|
||||
k, v = _expand_kv_heads(k, v, q.size(1))
|
||||
attn_mask, resolved_scale = _build_attn_mask(q, k, mask, causal_offset, scale)
|
||||
out = F.scaled_dot_product_attention(
|
||||
q, k, v, attn_mask=attn_mask, is_causal=False, scale=resolved_scale
|
||||
)
|
||||
# Restore Q's original layout
|
||||
if q_layout == 1:
|
||||
out = out.transpose(1, 2)
|
||||
return out
|
||||
|
||||
|
||||
def _gather_kv_from_pages(
|
||||
page_table: torch.Tensor,
|
||||
k_cache: torch.Tensor,
|
||||
v_cache: torch.Tensor,
|
||||
page_size: int,
|
||||
kv_len: int,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Gather contiguous K/V from paged cache for torch SDPA fallback.
|
||||
|
||||
Shapes:
|
||||
page_table : [batch, max_pages] (int64)
|
||||
k_cache : [n_pages, page_size, n_kv_heads, head_dim]
|
||||
v_cache : same as k_cache
|
||||
Returns:
|
||||
k, v : [batch, kv_len, n_kv_heads, head_dim] (b l h d)
|
||||
"""
|
||||
batch, max_pages = page_table.shape
|
||||
_, ps, n_kv_heads, head_dim = k_cache.shape
|
||||
if ps != page_size:
|
||||
raise ValueError(f"k_cache page_size mismatch: {ps} vs {page_size}")
|
||||
|
||||
# Vectorized gather: build physical page + offset indices, then advanced-index
|
||||
positions = torch.arange(kv_len, device=page_table.device)
|
||||
logical_pages = positions // page_size # [kv_len]
|
||||
page_offsets = positions % page_size # [kv_len]
|
||||
|
||||
phys_pages = page_table[:, logical_pages] # [batch, kv_len]
|
||||
# k_cache[phys_pages, page_offsets] → [batch, kv_len, n_kv_heads, head_dim] (b l h d)
|
||||
k = k_cache[phys_pages, page_offsets]
|
||||
v = v_cache[phys_pages, page_offsets]
|
||||
return k, v
|
||||
|
||||
|
||||
def attn_decode(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
mask: torch.Tensor | None = None,
|
||||
causal_offset: int = -1,
|
||||
scale: float = 0.0,
|
||||
layout: str = "bhld",
|
||||
) -> torch.Tensor:
|
||||
li = _parse_layout(layout)
|
||||
if _available["attn_decode"]:
|
||||
return _modules["attn_decode"].attn_decode(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
mask=mask,
|
||||
causal_offset=causal_offset,
|
||||
scale=scale,
|
||||
layout=li,
|
||||
)
|
||||
return _torch_fallback(q, k, v, mask, causal_offset, scale, q_layout=li)
|
||||
|
||||
|
||||
def attn_prefill(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
mask: torch.Tensor | None = None,
|
||||
causal_offset: int = -1,
|
||||
scale: float = 0.0,
|
||||
layout: str = "bhld",
|
||||
) -> torch.Tensor:
|
||||
li = _parse_layout(layout)
|
||||
if _available["attn_prefill"]:
|
||||
return _modules["attn_prefill"].attn_prefill(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
mask=mask,
|
||||
causal_offset=causal_offset,
|
||||
scale=scale,
|
||||
layout=li,
|
||||
)
|
||||
return _torch_fallback(q, k, v, mask, causal_offset, scale, q_layout=li)
|
||||
|
||||
|
||||
def attn_paged_decode(
|
||||
q: torch.Tensor,
|
||||
page_table: torch.Tensor,
|
||||
k_cache: torch.Tensor,
|
||||
v_cache: torch.Tensor,
|
||||
page_size: int,
|
||||
kv_len: int,
|
||||
mask: torch.Tensor | None = None,
|
||||
causal_offset: int = -1,
|
||||
scale: float = 0.0,
|
||||
layout: str = "bhld",
|
||||
) -> torch.Tensor:
|
||||
li = _parse_layout(layout)
|
||||
if _available["attn_paged_decode"]:
|
||||
return _modules["attn_paged_decode"].attn_paged_decode(
|
||||
q,
|
||||
page_table,
|
||||
k_cache,
|
||||
v_cache,
|
||||
page_size,
|
||||
kv_len,
|
||||
mask=mask,
|
||||
causal_offset=causal_offset,
|
||||
scale=scale,
|
||||
layout=li,
|
||||
)
|
||||
# Gathered K/V are always b l h d
|
||||
k, v = _gather_kv_from_pages(page_table, k_cache, v_cache, page_size, kv_len)
|
||||
return _torch_fallback(
|
||||
q, k, v, mask, causal_offset, scale, q_layout=li, kv_layout=1
|
||||
)
|
||||
|
||||
|
||||
def attention(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
mask: torch.Tensor | None = None,
|
||||
causal_offset: int = -1,
|
||||
scale: float = 0.0,
|
||||
layout: str = "bhld",
|
||||
) -> torch.Tensor:
|
||||
"""Dispatch to decode or prefill attention based on the query length.
|
||||
|
||||
A query length of one is the decode case; longer queries use prefill.
|
||||
The paged-cache decode path cannot be selected here because its page-table
|
||||
arguments are not part of this interface.
|
||||
"""
|
||||
li = _parse_layout(layout)
|
||||
|
||||
if q.ndim not in (2, 3, 4) or k.ndim != q.ndim or v.ndim != q.ndim:
|
||||
raise ValueError(
|
||||
"q, k, and v must all have the same rank in {2, 3, 4}, "
|
||||
f"got {q.ndim}D, {k.ndim}D, {v.ndim}D"
|
||||
)
|
||||
if k.shape != v.shape:
|
||||
raise ValueError(
|
||||
f"k and v must have the same shape, got {k.shape} and {v.shape}"
|
||||
)
|
||||
|
||||
original_ndim = q.ndim
|
||||
if original_ndim == 2:
|
||||
# [L, D] -> [1, 1, L, D] or [1, L, 1, D]
|
||||
q = q.unsqueeze(0).unsqueeze(1 if li == 0 else 2)
|
||||
k = k.unsqueeze(0).unsqueeze(1 if li == 0 else 2)
|
||||
v = v.unsqueeze(0).unsqueeze(1 if li == 0 else 2)
|
||||
elif original_ndim == 3:
|
||||
# [B, L, D] -> single-head 4D input.
|
||||
q = q.unsqueeze(1 if li == 0 else 2)
|
||||
k = k.unsqueeze(1 if li == 0 else 2)
|
||||
v = v.unsqueeze(1 if li == 0 else 2)
|
||||
|
||||
q_len = q.size(2 if li == 0 else 1)
|
||||
if q_len == 1:
|
||||
out = attn_decode(q, k, v, mask, causal_offset, scale, layout)
|
||||
else:
|
||||
out = attn_prefill(q, k, v, mask, causal_offset, scale, layout)
|
||||
|
||||
if original_ndim == 2:
|
||||
return out.squeeze(0).squeeze(0 if li == 0 else 1)
|
||||
if original_ndim == 3:
|
||||
return out.squeeze(1 if li == 0 else 2)
|
||||
return out
|
||||
+41
-35
@@ -13,41 +13,63 @@ from typing import (
|
||||
Type,
|
||||
TypeVar,
|
||||
Union,
|
||||
get_args,
|
||||
get_origin,
|
||||
)
|
||||
from typing import get_args as _get_args
|
||||
from typing import get_origin as _get_origin
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
def _resolve_type(
|
||||
def _resolve_base_type(
|
||||
arg: Union[Type, str, ForwardRef], factory_cls: type
|
||||
) -> Optional[Type]:
|
||||
"""Resolve a generic type-arg (str forward-ref, ForwardRef, or class)."""
|
||||
if not isinstance(arg, (str, ForwardRef)):
|
||||
"""Resolve the generic type-arg T to a concrete class.
|
||||
|
||||
- Concrete class (``BaseFactory[MyBase]``): returned directly.
|
||||
- Forward reference (``BaseFactory["MyBase"]``): ``Base["X"]``
|
||||
produces a ``ForwardRef("X")`` at class-creation time. We
|
||||
extract the name and evaluate it in the factory module's
|
||||
global namespace — the same mechanism ``typing.get_type_hints``
|
||||
uses internally.
|
||||
"""
|
||||
if isinstance(arg, type):
|
||||
return arg
|
||||
|
||||
name = arg if isinstance(arg, str) else arg.__forward_arg__
|
||||
if name == factory_cls.__name__:
|
||||
return factory_cls
|
||||
if isinstance(arg, str):
|
||||
name = arg
|
||||
elif isinstance(arg, ForwardRef):
|
||||
name = arg.__forward_arg__
|
||||
else:
|
||||
return None
|
||||
|
||||
mod = sys.modules.get(factory_cls.__module__)
|
||||
if mod is None:
|
||||
return None
|
||||
ns = vars(mod)
|
||||
try:
|
||||
return eval(name, vars(mod)) # noqa: S307
|
||||
except NameError:
|
||||
return None
|
||||
|
||||
if isinstance(arg, ForwardRef):
|
||||
return arg._evaluate(ns, None, recursive_guard=frozenset())
|
||||
|
||||
return ns.get(name)
|
||||
def _validate_component(component_cls: Type, base: Optional[Type]) -> None:
|
||||
"""Validate that *component_cls* inherits from *base*.
|
||||
|
||||
No-op when *base* is ``None`` (e.g. forward-ref resolution failed).
|
||||
"""
|
||||
if base is not None and not issubclass(component_cls, base):
|
||||
raise TypeError(f"{component_cls.__name__} must inherit from {base.__name__}")
|
||||
|
||||
|
||||
class BaseFactory(ABC, Generic[T]):
|
||||
"""Generic factory with decorator-based component registration.
|
||||
"""Generic factory with decorator-based registration.
|
||||
|
||||
Create a factory by subclassing with the desired base type::
|
||||
|
||||
class MyFactory(BaseFactory[MyBase]):
|
||||
pass
|
||||
|
||||
Register components with the ``register`` decorator::
|
||||
|
||||
@MyFactory.register("custom")
|
||||
class CustomComponent(MyBase):
|
||||
...
|
||||
@@ -64,13 +86,10 @@ class BaseFactory(ABC, Generic[T]):
|
||||
def __init_subclass__(cls, **kwargs):
|
||||
super().__init_subclass__(**kwargs)
|
||||
for orig_base in getattr(cls, "__orig_bases__", ()):
|
||||
if _get_origin(orig_base) is BaseFactory:
|
||||
(arg,) = _get_args(orig_base)
|
||||
if get_origin(orig_base) is BaseFactory:
|
||||
(arg,) = get_args(orig_base)
|
||||
cls._entries = {}
|
||||
try:
|
||||
cls._component_base = _resolve_type(arg, cls)
|
||||
except Exception:
|
||||
cls._component_base = None
|
||||
cls._component_base = _resolve_base_type(arg, cls)
|
||||
return
|
||||
|
||||
@classmethod
|
||||
@@ -82,7 +101,7 @@ class BaseFactory(ABC, Generic[T]):
|
||||
"""
|
||||
|
||||
def decorator(component_cls: Type[T]) -> Type[T]:
|
||||
cls._validate_component(component_cls)
|
||||
_validate_component(component_cls, cls._component_base)
|
||||
if name in cls._entries:
|
||||
raise ValueError(f"Component '{name}' is already registered")
|
||||
cls._entries[name] = component_cls
|
||||
@@ -95,12 +114,11 @@ class BaseFactory(ABC, Generic[T]):
|
||||
"""Create a component instance by name, filtering kwargs to match
|
||||
the component's ``__init__`` signature.
|
||||
"""
|
||||
entry = cls._entries.get(name)
|
||||
if entry is None:
|
||||
component_cls = cls._entries.get(name)
|
||||
if component_cls is None:
|
||||
raise ValueError(
|
||||
f"Unknown component: '{name}'. Supported types: {sorted(cls._entries)}"
|
||||
)
|
||||
component_cls = entry
|
||||
sig = inspect.signature(component_cls.__init__)
|
||||
has_var_kwargs = any(
|
||||
p.kind == inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values()
|
||||
@@ -114,18 +132,6 @@ class BaseFactory(ABC, Generic[T]):
|
||||
kwargs = {k: v for k, v in kwargs.items() if k in valid}
|
||||
return component_cls(*args, **kwargs)
|
||||
|
||||
@classmethod
|
||||
def _validate_component(cls, component_cls: Type[T]):
|
||||
"""Validate the decorated class inherits from the factory's base type.
|
||||
|
||||
Override for custom validation beyond ``issubclass``.
|
||||
"""
|
||||
base = cls._component_base
|
||||
if base is not None and not issubclass(component_cls, base):
|
||||
raise TypeError(
|
||||
f"{component_cls.__name__} must inherit from {base.__name__}"
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_component_class(cls, name: str) -> Type[T]:
|
||||
"""Get the registered component class without instantiating it."""
|
||||
|
||||
@@ -30,21 +30,16 @@ from astrai.inference.api.openai import OpenAIResponseBuilder
|
||||
from astrai.inference.core import (
|
||||
STOP,
|
||||
Allocator,
|
||||
CacheView,
|
||||
ContiguousCache,
|
||||
ContiguousCacheView,
|
||||
Executor,
|
||||
InferenceScheduler,
|
||||
KVCache,
|
||||
PageCache,
|
||||
PageCacheView,
|
||||
KVStorage,
|
||||
PagePool,
|
||||
PrefixCache,
|
||||
Storage,
|
||||
ReqToTokenPool,
|
||||
Task,
|
||||
TaskManager,
|
||||
TaskStatus,
|
||||
TaskTable,
|
||||
page_hash,
|
||||
)
|
||||
from astrai.inference.engine import GenerationRequest, InferenceEngine
|
||||
@@ -68,16 +63,11 @@ __all__ = [
|
||||
"TaskManager",
|
||||
"TaskStatus",
|
||||
"Allocator",
|
||||
"CacheView",
|
||||
"KVCache",
|
||||
"ContiguousCache",
|
||||
"ContiguousCacheView",
|
||||
"PageCache",
|
||||
"PageCacheView",
|
||||
"KVStorage",
|
||||
"PagePool",
|
||||
"PrefixCache",
|
||||
"Storage",
|
||||
"TaskTable",
|
||||
"ReqToTokenPool",
|
||||
"page_hash",
|
||||
"sample",
|
||||
"BaseSamplingStrategy",
|
||||
|
||||
@@ -110,6 +110,7 @@ def _create_engine(
|
||||
device: str = "cuda",
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
max_batch_size: int = 16,
|
||||
max_seq_len: Optional[int] = None,
|
||||
) -> InferenceEngine:
|
||||
if not param_path.exists():
|
||||
raise FileNotFoundError(f"Parameter directory not found: {param_path}")
|
||||
@@ -123,6 +124,7 @@ def _create_engine(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=max_batch_size,
|
||||
max_seq_len=max_seq_len,
|
||||
)
|
||||
logger.info(f"Inference engine initialized with max_batch_size={max_batch_size}")
|
||||
return engine
|
||||
@@ -186,6 +188,7 @@ def run_server(
|
||||
device: str = "cuda",
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
max_batch_size: int = 16,
|
||||
max_seq_len: Optional[int] = None,
|
||||
):
|
||||
app = get_app()
|
||||
app.state.server_config = {
|
||||
@@ -193,6 +196,7 @@ def run_server(
|
||||
"dtype": dtype,
|
||||
"param_path": param_path,
|
||||
"max_batch_size": max_batch_size,
|
||||
"max_seq_len": max_seq_len,
|
||||
}
|
||||
uvicorn.run(
|
||||
app,
|
||||
|
||||
@@ -22,13 +22,10 @@ class BaseToolParser(ABC):
|
||||
Maintains streaming state internally so that each call to :meth:`feed`
|
||||
can diff against previously emitted content.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
tools : list of dict, optional
|
||||
Tool definitions from the request.
|
||||
tool_choice : str
|
||||
``"auto"`` / ``"required"`` / ``"none"`` or a named tool choice
|
||||
dict.
|
||||
Args:
|
||||
tools (list of dict, optional): Tool definitions from the request.
|
||||
tool_choice (str): ``"auto"`` / ``"required"`` / ``"none"`` or a named
|
||||
tool choice dict.
|
||||
"""
|
||||
|
||||
def __init__(self, tools: Optional[List[Dict]] = None, tool_choice: str = "auto"):
|
||||
@@ -51,14 +48,12 @@ class BaseToolParser(ABC):
|
||||
|
||||
Returns an empty list when nothing new should be emitted.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
body : str
|
||||
The complete accumulated generated text so far.
|
||||
current_token_ids : list of int, optional
|
||||
All token IDs decoded into *body* (cumulative).
|
||||
delta_token_ids : list of int, optional
|
||||
Only the token IDs for this chunk.
|
||||
Args:
|
||||
body (str): The complete accumulated generated text so far.
|
||||
current_token_ids (list of int, optional): All token IDs decoded
|
||||
into *body* (cumulative).
|
||||
delta_token_ids (list of int, optional): Only the token IDs for
|
||||
this chunk.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
|
||||
@@ -2,16 +2,11 @@
|
||||
|
||||
from astrai.inference.core.cache import (
|
||||
Allocator,
|
||||
CacheView,
|
||||
ContiguousCache,
|
||||
ContiguousCacheView,
|
||||
KVCache,
|
||||
PageCache,
|
||||
PageCacheView,
|
||||
KVStorage,
|
||||
PagePool,
|
||||
PrefixCache,
|
||||
Storage,
|
||||
TaskTable,
|
||||
ReqToTokenPool,
|
||||
page_hash,
|
||||
)
|
||||
from astrai.inference.core.executor import Executor
|
||||
@@ -20,16 +15,11 @@ from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
|
||||
|
||||
__all__ = [
|
||||
"Allocator",
|
||||
"CacheView",
|
||||
"KVCache",
|
||||
"ContiguousCache",
|
||||
"ContiguousCacheView",
|
||||
"PageCache",
|
||||
"PageCacheView",
|
||||
"KVStorage",
|
||||
"PagePool",
|
||||
"PrefixCache",
|
||||
"Storage",
|
||||
"TaskTable",
|
||||
"ReqToTokenPool",
|
||||
"page_hash",
|
||||
"Executor",
|
||||
"InferenceScheduler",
|
||||
|
||||
+315
-357
@@ -1,7 +1,21 @@
|
||||
"""KV cache architecture: three-layer separation (SGLang-inspired).
|
||||
|
||||
Layer 1 — KVStorage: flat token-level K/V buffers [n_layers, size, H, D]
|
||||
Layer 2 — ReqToTokenPool: index table [req_idx, pos] → physical token slot
|
||||
Layer 3 — Allocator: slot/page allocation with ref-counting and LRU
|
||||
|
||||
PagePool orchestrates all three plus PrefixCache (content addressing).
|
||||
KVCache is a pure dataclass passed to the model for direct buffer access.
|
||||
|
||||
Two modes:
|
||||
- contiguous (default): pre-allocated per-request blocks, no dynamic alloc
|
||||
- paged: shared pool with on-demand allocation, prefix caching support
|
||||
"""
|
||||
|
||||
import threading
|
||||
from abc import ABC, abstractmethod
|
||||
from collections import OrderedDict
|
||||
from typing import Callable, Dict, List, Optional, Tuple
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable, Dict, List, Optional
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
@@ -108,418 +122,362 @@ class PrefixCache:
|
||||
self._hash_to_page[h] = page_idx
|
||||
|
||||
|
||||
class PagePool:
|
||||
"""Orchestrates allocator (page management) and PrefixCache (content addressing)."""
|
||||
class ReqToTokenPool:
|
||||
"""Maps [req_idx, pos] -> physical token slot in KV storage.
|
||||
|
||||
def __init__(self, allocator: Allocator, prefix: PrefixCache):
|
||||
self._alloc = allocator
|
||||
self._prefix = prefix
|
||||
self._alloc.on_evict = prefix.evict
|
||||
Each row is one request; each column is a sequence position. The value
|
||||
at [req_idx, pos] is the flat index into the KV storage buffers.
|
||||
"""
|
||||
|
||||
@property
|
||||
def allocator(self) -> Allocator:
|
||||
return self._alloc
|
||||
|
||||
@property
|
||||
def prefix(self) -> PrefixCache:
|
||||
return self._prefix
|
||||
|
||||
def alloc(self) -> int:
|
||||
return self._alloc.alloc()
|
||||
|
||||
def free(self, idx: int):
|
||||
keep = self._prefix.has_page(idx)
|
||||
self._alloc.free(idx, keep_cached=keep)
|
||||
if not keep:
|
||||
self._prefix.evict(idx)
|
||||
|
||||
def inc_ref(self, idx: int):
|
||||
self._alloc.inc_ref(idx)
|
||||
|
||||
def lookup(self, token_ids: List[int]) -> List[int]:
|
||||
hits = self._prefix.lookup(token_ids)
|
||||
for p in hits:
|
||||
self._alloc.touch(p)
|
||||
return hits
|
||||
|
||||
def record(self, page_idx: int, token_ids: List[int], logical_page_idx: int):
|
||||
self._prefix.record(page_idx, token_ids, logical_page_idx)
|
||||
|
||||
|
||||
class TaskTable:
|
||||
"""Maps task_ids to page tables and cached token counts."""
|
||||
|
||||
def __init__(self, page_size: int):
|
||||
self._page_size = page_size
|
||||
self._pages: Dict[str, List[int]] = {}
|
||||
self._cached: Dict[str, int] = {}
|
||||
def __init__(self, size: int, max_context_len: int, device: torch.device):
|
||||
self.size = size
|
||||
self.max_context_len = max_context_len
|
||||
self.req_to_token = torch.zeros(
|
||||
(size, max_context_len), dtype=torch.long, device=device
|
||||
)
|
||||
self.free_slots = list(range(size))
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def set(self, task_id: str, page_table: List[int], cached: int):
|
||||
def alloc(self, num_reqs: int) -> Optional[List[int]]:
|
||||
with self._lock:
|
||||
self._pages[task_id] = page_table
|
||||
self._cached[task_id] = cached
|
||||
if num_reqs > len(self.free_slots):
|
||||
return None
|
||||
slots = self.free_slots[:num_reqs]
|
||||
self.free_slots = self.free_slots[num_reqs:]
|
||||
return slots
|
||||
|
||||
def get(self, task_id: str) -> List[int]:
|
||||
def free(self, req_indices: List[int]):
|
||||
with self._lock:
|
||||
return self._pages.get(task_id, [])
|
||||
self.free_slots.extend(req_indices)
|
||||
|
||||
def get_cached(self, task_id: str) -> int:
|
||||
with self._lock:
|
||||
return self._cached.get(task_id, 0)
|
||||
|
||||
def pop(self, task_id: str) -> Tuple[List[int], int]:
|
||||
with self._lock:
|
||||
pages = self._pages.pop(task_id, [])
|
||||
cached = self._cached.pop(task_id, 0)
|
||||
return pages, cached
|
||||
|
||||
def get_ref(self, task_id: str) -> List[int]:
|
||||
with self._lock:
|
||||
return self._pages.setdefault(task_id, [])
|
||||
|
||||
def table_tensor(self, task_ids: List[str], device: torch.device) -> Tensor:
|
||||
with self._lock:
|
||||
states = [self._pages.get(tid, []) for tid in task_ids]
|
||||
max_pages = max((len(s) for s in states), default=0)
|
||||
rows = [s + [-1] * (max_pages - len(s)) for s in states]
|
||||
return torch.tensor(rows, dtype=torch.long, device=device)
|
||||
def write(self, indices, values):
|
||||
self.req_to_token[indices] = values
|
||||
|
||||
|
||||
class Storage:
|
||||
"""KV-cache tensor storage with paged write/gather."""
|
||||
class KVStorage:
|
||||
"""Token-level KV cache storage.
|
||||
|
||||
Buffers: [n_layers, size, n_kv_heads, head_dim]. Each token occupies
|
||||
one slot indexed by ReqToTokenPool.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size: int,
|
||||
n_layers: int,
|
||||
n_pages: int,
|
||||
page_size: int,
|
||||
n_kv_heads: int,
|
||||
head_dim: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
):
|
||||
self.page_size = page_size
|
||||
self.k_cache = torch.empty(
|
||||
(n_layers, n_pages, page_size, n_kv_heads, head_dim),
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
self.size = size
|
||||
self.k_buffer = torch.empty(
|
||||
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype
|
||||
)
|
||||
self.v_cache = torch.empty(
|
||||
(n_layers, n_pages, page_size, n_kv_heads, head_dim),
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
self.v_buffer = torch.empty(
|
||||
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype
|
||||
)
|
||||
|
||||
def write(
|
||||
self,
|
||||
layer_id: int,
|
||||
page_table: Tensor,
|
||||
start_pos: int,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
):
|
||||
seq_len = k.size(1)
|
||||
if seq_len == 0:
|
||||
return
|
||||
page_size = self.page_size
|
||||
written = 0
|
||||
first_page = start_pos // page_size
|
||||
last_page = (start_pos + seq_len - 1) // page_size
|
||||
for pi in range(first_page, last_page + 1):
|
||||
phys_pages = page_table[:, pi]
|
||||
page_start = pi * page_size
|
||||
write_start = max(page_start, start_pos)
|
||||
write_end = min(page_start + page_size, start_pos + seq_len)
|
||||
offset = write_start - page_start
|
||||
chunk = write_end - write_start
|
||||
valid = phys_pages >= 0
|
||||
if not valid.all():
|
||||
if valid.any():
|
||||
valid_pages = phys_pages[valid]
|
||||
self.k_cache[layer_id, valid_pages, offset : offset + chunk] = k[
|
||||
valid, written : written + chunk
|
||||
]
|
||||
self.v_cache[layer_id, valid_pages, offset : offset + chunk] = v[
|
||||
valid, written : written + chunk
|
||||
]
|
||||
written += chunk
|
||||
continue
|
||||
self.k_cache[layer_id, phys_pages, offset : offset + chunk] = k[
|
||||
:, written : written + chunk
|
||||
]
|
||||
self.v_cache[layer_id, phys_pages, offset : offset + chunk] = v[
|
||||
:, written : written + chunk
|
||||
]
|
||||
written += chunk
|
||||
def get_key_buffer(self, layer_id: int) -> Tensor:
|
||||
return self.k_buffer[layer_id]
|
||||
|
||||
def gather(
|
||||
self, layer_id: int, page_table: Tensor, total_len: int
|
||||
) -> Tuple[Tensor, Tensor]:
|
||||
safe = page_table.clamp(min=0)
|
||||
k = self.k_cache[layer_id, safe]
|
||||
v = self.v_cache[layer_id, safe]
|
||||
k = k.flatten(1, 2)
|
||||
v = v.flatten(1, 2)
|
||||
if (page_table < 0).any():
|
||||
invalid = (
|
||||
(page_table < 0)
|
||||
.unsqueeze(-1)
|
||||
.expand(-1, -1, self.page_size)
|
||||
.flatten(1, 2)
|
||||
)
|
||||
invalid = invalid[:, :, None, None].expand_as(k)
|
||||
k = k.masked_fill(invalid, 0.0)
|
||||
v = v.masked_fill(invalid, 0.0)
|
||||
k = k[:, :total_len]
|
||||
v = v[:, :total_len]
|
||||
return k, v
|
||||
def get_value_buffer(self, layer_id: int) -> Tensor:
|
||||
return self.v_buffer[layer_id]
|
||||
|
||||
def set_kv_buffer(self, layer_id: int, loc: Tensor, k: Tensor, v: Tensor) -> None:
|
||||
self.k_buffer[layer_id, loc] = k
|
||||
self.v_buffer[layer_id, loc] = v
|
||||
|
||||
|
||||
class CacheView(ABC):
|
||||
"""Abstract view passed to attention layers for KV-cache I/O."""
|
||||
@dataclass
|
||||
class KVCache:
|
||||
"""Pure data struct passed to model for KV cache I/O.
|
||||
|
||||
@abstractmethod
|
||||
def write(self, layer_id: int, k: Tensor, v: Tensor): ...
|
||||
The attention layer does raw buffer indexing — no methods, no abstraction.
|
||||
|
||||
@abstractmethod
|
||||
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]: ...
|
||||
Attributes:
|
||||
k_buffer: [n_layers, size, n_kv_heads, head_dim]
|
||||
v_buffer: [n_layers, size, n_kv_heads, head_dim]
|
||||
req_to_token: [num_reqs, max_ctx_len] — index table
|
||||
req_pool_indices: [batch_size] — row indices into req_to_token
|
||||
seq_lens: [batch_size] — per-request total sequence lengths
|
||||
out_cache_loc: [batch, new_seq_len] or [batch, 1] — write indices
|
||||
max_len: max(seq_lens) as Python int — avoids GPU sync in decode
|
||||
"""
|
||||
|
||||
k_buffer: Tensor
|
||||
v_buffer: Tensor
|
||||
req_to_token: Tensor
|
||||
req_pool_indices: Tensor
|
||||
seq_lens: Tensor
|
||||
out_cache_loc: Tensor
|
||||
max_len: int = 0
|
||||
|
||||
|
||||
class KVCache(ABC):
|
||||
"""Abstract KV-cache facade for scheduler/executor."""
|
||||
class PagePool:
|
||||
"""Top-level KV cache manager.
|
||||
|
||||
@abstractmethod
|
||||
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool: ...
|
||||
Combines KVStorage + ReqToTokenPool + Allocator + PrefixCache.
|
||||
|
||||
@abstractmethod
|
||||
def task_free(self, task_id: str): ...
|
||||
|
||||
@abstractmethod
|
||||
def task_extend(self, task_id: str, pos: int) -> bool: ...
|
||||
|
||||
@abstractmethod
|
||||
def bind_tasks(
|
||||
self,
|
||||
task_ids: List[str],
|
||||
total_len: int,
|
||||
device: torch.device,
|
||||
write_positions: Optional[Tensor] = None,
|
||||
) -> CacheView: ...
|
||||
|
||||
def task_cached(self, task_id: str) -> int:
|
||||
return 0
|
||||
|
||||
def task_record_hashes(
|
||||
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
||||
): ...
|
||||
|
||||
|
||||
class PageCacheView(CacheView):
|
||||
"""Bundles Storage + page_table + total_len for attention layers."""
|
||||
|
||||
def __init__(self, storage: Storage, page_table: Tensor, total_len: int = 0):
|
||||
self._storage = storage
|
||||
self._page_table = page_table
|
||||
self._total_len = total_len
|
||||
|
||||
def write(self, layer_id: int, k: Tensor, v: Tensor):
|
||||
start_pos = self._total_len - k.size(1)
|
||||
self._storage.write(layer_id, self._page_table, start_pos, k, v)
|
||||
|
||||
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
|
||||
return self._storage.gather(layer_id, self._page_table, self._total_len)
|
||||
|
||||
|
||||
class PageCache(KVCache):
|
||||
"""Paged KV-cache with prefix sharing."""
|
||||
Args:
|
||||
n_layers: Number of transformer layers.
|
||||
n_kv_heads: Number of KV attention heads.
|
||||
head_dim: Dimension per head.
|
||||
max_batch_size: Maximum concurrent requests.
|
||||
max_seq_len: Maximum sequence length per request.
|
||||
device, dtype: Tensor device and dtype.
|
||||
page_size: Page size for paged mode (1 = token-level).
|
||||
n_tokens: Total token slots for paged mode. None = contiguous mode
|
||||
(pre-allocates max_batch_size * max_seq_len).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
n_layers: int,
|
||||
n_pages: int,
|
||||
page_size: int,
|
||||
n_kv_heads: int,
|
||||
head_dim: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
):
|
||||
self.page_size = page_size
|
||||
self._pool = PagePool(Allocator(n_pages), PrefixCache(page_size))
|
||||
self._table = TaskTable(page_size)
|
||||
self._storage = Storage(
|
||||
n_layers, n_pages, page_size, n_kv_heads, head_dim, device, dtype
|
||||
)
|
||||
|
||||
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
||||
hits = self._pool.lookup(prompt_ids)
|
||||
cached = len(hits) * self.page_size
|
||||
for p in hits:
|
||||
self._pool.inc_ref(p)
|
||||
|
||||
remaining = len(prompt_ids) - cached
|
||||
n_new = (
|
||||
(remaining + self.page_size - 1) // self.page_size if remaining > 0 else 0
|
||||
)
|
||||
new_pages: List[int] = []
|
||||
if n_new > 0:
|
||||
for _ in range(n_new):
|
||||
p = self._pool.alloc()
|
||||
if p < 0:
|
||||
for hp in hits:
|
||||
self._pool.free(hp)
|
||||
for np in new_pages:
|
||||
self._pool.free(np)
|
||||
return False
|
||||
new_pages.append(p)
|
||||
|
||||
self._table.set(task_id, hits + new_pages, cached)
|
||||
return True
|
||||
|
||||
def task_free(self, task_id: str):
|
||||
page_table, _ = self._table.pop(task_id)
|
||||
for idx in page_table:
|
||||
self._pool.free(idx)
|
||||
|
||||
def task_extend(self, task_id: str, pos: int) -> bool:
|
||||
page_table = self._table.get(task_id)
|
||||
needed = (pos + 1 + self.page_size - 1) // self.page_size
|
||||
while len(page_table) < needed:
|
||||
p = self._pool.alloc()
|
||||
if p < 0:
|
||||
return False
|
||||
page_table.append(p)
|
||||
return True
|
||||
|
||||
def task_cached(self, task_id: str) -> int:
|
||||
return self._table.get_cached(task_id)
|
||||
|
||||
def task_record_hashes(
|
||||
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
||||
):
|
||||
page_table = self._table.get(task_id)
|
||||
full_pages = len(prompt_ids) // self.page_size
|
||||
for i in range(start_logical_page, full_pages):
|
||||
self._pool.record(page_table[i], prompt_ids, i)
|
||||
|
||||
def bind_tasks(
|
||||
self,
|
||||
task_ids: List[str],
|
||||
total_len: int,
|
||||
device: torch.device,
|
||||
write_positions: Optional[Tensor] = None,
|
||||
) -> PageCacheView:
|
||||
page_table = self._table.table_tensor(task_ids, device)
|
||||
return PageCacheView(self._storage, page_table, total_len)
|
||||
|
||||
|
||||
class ContiguousCacheView(CacheView):
|
||||
"""Contiguous KV-cache view for attention layers."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
cache: "ContiguousCache",
|
||||
batch_indices: Tensor,
|
||||
total_len: int = 0,
|
||||
write_positions: Optional[Tensor] = None,
|
||||
):
|
||||
self._cache = cache
|
||||
self._batch_indices = batch_indices
|
||||
self._total_len = total_len
|
||||
self._write_positions = write_positions
|
||||
|
||||
def write(self, layer_id: int, k: Tensor, v: Tensor):
|
||||
seq_len = k.size(1)
|
||||
indices = self._batch_indices
|
||||
if self._write_positions is not None and seq_len == 1:
|
||||
pos = self._write_positions
|
||||
self._cache.k[layer_id, indices, pos] = k.squeeze(1)
|
||||
self._cache.v[layer_id, indices, pos] = v.squeeze(1)
|
||||
else:
|
||||
start_pos = self._total_len - seq_len
|
||||
self._cache.k[layer_id, indices, start_pos : start_pos + seq_len] = k
|
||||
self._cache.v[layer_id, indices, start_pos : start_pos + seq_len] = v
|
||||
|
||||
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
|
||||
max_len = self._total_len
|
||||
indices = self._batch_indices
|
||||
k = self._cache.k[layer_id, indices, :max_len]
|
||||
v = self._cache.v[layer_id, indices, :max_len]
|
||||
return k, v
|
||||
|
||||
|
||||
class ContiguousCache(KVCache):
|
||||
"""Contiguous per-slot KV cache (default implementation)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
n_layers: int,
|
||||
max_batch_size: int,
|
||||
max_seq_len: int,
|
||||
n_kv_heads: int,
|
||||
head_dim: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
page_size: int = 1,
|
||||
n_tokens: Optional[int] = None,
|
||||
):
|
||||
self.page_size = page_size
|
||||
self.max_batch_size = max_batch_size
|
||||
self.max_seq_len = max_seq_len
|
||||
self.k = torch.zeros(
|
||||
n_layers,
|
||||
max_batch_size,
|
||||
max_seq_len,
|
||||
n_kv_heads,
|
||||
head_dim,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
self.device = device
|
||||
self.dtype = dtype
|
||||
self.n_layers = n_layers
|
||||
self.n_kv_heads = n_kv_heads
|
||||
self.head_dim = head_dim
|
||||
|
||||
self.contiguous = n_tokens is None
|
||||
if self.contiguous:
|
||||
self.n_tokens = max_batch_size * max_seq_len
|
||||
else:
|
||||
self.n_tokens = n_tokens
|
||||
|
||||
self._storage = KVStorage(
|
||||
self.n_tokens, n_layers, n_kv_heads, head_dim, device, dtype
|
||||
)
|
||||
self.v = torch.zeros(
|
||||
n_layers,
|
||||
max_batch_size,
|
||||
max_seq_len,
|
||||
n_kv_heads,
|
||||
head_dim,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
self._slot_len: Dict[int, int] = {}
|
||||
self._task_slot: Dict[str, int] = {}
|
||||
self._free_slots = list(range(max_batch_size))
|
||||
self._device = device
|
||||
self._req_pool = ReqToTokenPool(max_batch_size, max_seq_len, device)
|
||||
|
||||
if self.contiguous:
|
||||
for i in range(max_batch_size):
|
||||
self._req_pool.req_to_token[i] = torch.arange(
|
||||
i * max_seq_len, (i + 1) * max_seq_len, device=device
|
||||
)
|
||||
self._alloc: Optional[Allocator] = None
|
||||
self._prefix: Optional[PrefixCache] = None
|
||||
else:
|
||||
n_pages = self.n_tokens // page_size
|
||||
self._alloc = Allocator(n_pages)
|
||||
self._prefix = PrefixCache(page_size) if page_size > 1 else None
|
||||
if self._prefix is not None:
|
||||
self._alloc.on_evict = self._prefix.evict
|
||||
|
||||
self._task_req: Dict[str, int] = {}
|
||||
self._task_len: Dict[int, int] = {}
|
||||
self._task_cached: Dict[str, int] = {}
|
||||
self._task_slots: Dict[str, List[int]] = {}
|
||||
self._task_pages: Dict[str, List[int]] = {}
|
||||
self._lock = threading.Lock()
|
||||
|
||||
# ---- task lifecycle ----
|
||||
|
||||
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
||||
if not self._free_slots:
|
||||
req_slots = self._req_pool.alloc(1)
|
||||
if req_slots is None:
|
||||
return False
|
||||
slot = self._free_slots.pop(0)
|
||||
self._task_slot[task_id] = slot
|
||||
self._slot_len[slot] = 0
|
||||
req_idx = req_slots[0]
|
||||
self._task_req[task_id] = req_idx
|
||||
|
||||
if self.contiguous:
|
||||
self._task_len[req_idx] = len(prompt_ids)
|
||||
self._task_cached[task_id] = 0
|
||||
return True
|
||||
|
||||
n_tokens_needed = len(prompt_ids)
|
||||
cached = 0
|
||||
|
||||
if self._prefix is not None:
|
||||
hits = self._prefix.lookup(prompt_ids)
|
||||
cached = len(hits) * self.page_size
|
||||
for p in hits:
|
||||
self._alloc.inc_ref(p)
|
||||
self._task_pages[task_id] = list(hits)
|
||||
self._task_slots[task_id] = []
|
||||
else:
|
||||
self._task_pages[task_id] = []
|
||||
self._task_slots[task_id] = []
|
||||
|
||||
remaining = n_tokens_needed - cached
|
||||
if remaining > 0:
|
||||
if self.page_size == 1:
|
||||
slots = self._alloc_tokens(remaining)
|
||||
if slots is None:
|
||||
for p in self._task_pages[task_id]:
|
||||
self._alloc.free(p)
|
||||
self._req_pool.free([req_idx])
|
||||
del self._task_req[task_id]
|
||||
return False
|
||||
self._task_slots[task_id] = slots
|
||||
else:
|
||||
n_new_pages = (remaining + self.page_size - 1) // self.page_size
|
||||
new_pages = []
|
||||
for _ in range(n_new_pages):
|
||||
p = self._alloc.alloc()
|
||||
if p < 0:
|
||||
for hp in self._task_pages[task_id]:
|
||||
self._alloc.free(hp)
|
||||
for np_ in new_pages:
|
||||
self._alloc.free(np_)
|
||||
self._req_pool.free([req_idx])
|
||||
del self._task_req[task_id]
|
||||
return False
|
||||
new_pages.append(p)
|
||||
self._task_pages[task_id].extend(new_pages)
|
||||
|
||||
self._write_req_to_token(task_id, prompt_ids, cached)
|
||||
self._task_len[req_idx] = len(prompt_ids)
|
||||
self._task_cached[task_id] = cached
|
||||
return True
|
||||
|
||||
def task_free(self, task_id: str):
|
||||
slot = self._task_slot.pop(task_id, None)
|
||||
if slot is not None:
|
||||
self._slot_len.pop(slot, None)
|
||||
self._free_slots.append(slot)
|
||||
req_idx = self._task_req.pop(task_id, None)
|
||||
if req_idx is None:
|
||||
return
|
||||
self._task_len.pop(req_idx, None)
|
||||
self._task_cached.pop(task_id, None)
|
||||
|
||||
if not self.contiguous:
|
||||
if self._prefix is not None:
|
||||
for p in self._task_pages.get(task_id, []):
|
||||
keep = self._prefix.has_page(p)
|
||||
self._alloc.free(p, keep_cached=keep)
|
||||
if not keep:
|
||||
self._prefix.evict(p)
|
||||
else:
|
||||
for p in self._task_pages.get(task_id, []):
|
||||
self._alloc.free(p)
|
||||
self._task_pages.pop(task_id, None)
|
||||
self._task_slots.pop(task_id, None)
|
||||
|
||||
self._req_pool.free([req_idx])
|
||||
|
||||
def task_extend(self, task_id: str, pos: int) -> bool:
|
||||
return pos < self.max_seq_len
|
||||
req_idx = self._task_req.get(task_id)
|
||||
if req_idx is None:
|
||||
return False
|
||||
|
||||
if self.contiguous:
|
||||
return pos < self.max_seq_len
|
||||
|
||||
if self.page_size == 1:
|
||||
slots = self._alloc_tokens(1)
|
||||
if slots is None:
|
||||
return False
|
||||
self._task_slots.setdefault(task_id, []).extend(slots)
|
||||
self._req_pool.req_to_token[req_idx, pos] = slots[0]
|
||||
else:
|
||||
page_idx = pos // self.page_size
|
||||
existing = self._task_pages.get(task_id, [])
|
||||
if page_idx >= len(existing):
|
||||
p = self._alloc.alloc()
|
||||
if p < 0:
|
||||
return False
|
||||
existing.append(p)
|
||||
self._task_pages[task_id] = existing
|
||||
page_offset = pos % self.page_size
|
||||
page = existing[page_idx]
|
||||
token_slot = page * self.page_size + page_offset
|
||||
self._req_pool.req_to_token[req_idx, pos] = token_slot
|
||||
|
||||
self._task_len[req_idx] = pos + 1
|
||||
return True
|
||||
|
||||
def task_cached(self, task_id: str) -> int:
|
||||
slot = self._task_slot.get(task_id)
|
||||
if slot is None:
|
||||
return 0
|
||||
return self._slot_len.get(slot, 0)
|
||||
return self._task_cached.get(task_id, 0)
|
||||
|
||||
def task_record_hashes(
|
||||
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
||||
):
|
||||
if self._prefix is None or self.contiguous:
|
||||
return
|
||||
pages = self._task_pages.get(task_id, [])
|
||||
full_pages = len(prompt_ids) // self.page_size
|
||||
for i in range(start_logical_page, min(full_pages, len(pages))):
|
||||
self._prefix.record(pages[i], prompt_ids, i)
|
||||
|
||||
# ---- bind for forward ----
|
||||
|
||||
def bind_tasks(
|
||||
self,
|
||||
task_ids: List[str],
|
||||
total_len: int,
|
||||
seq_lens: List[int],
|
||||
device: torch.device,
|
||||
write_positions: Optional[Tensor] = None,
|
||||
) -> ContiguousCacheView:
|
||||
slots = [self._task_slot[tid] for tid in task_ids]
|
||||
batch_indices = torch.tensor(slots, dtype=torch.long, device=device)
|
||||
for slot in slots:
|
||||
if total_len > self._slot_len.get(slot, 0):
|
||||
self._slot_len[slot] = total_len
|
||||
return ContiguousCacheView(
|
||||
self, batch_indices, total_len, write_positions=write_positions
|
||||
start_pos: Optional[int] = None,
|
||||
) -> KVCache:
|
||||
req_indices = [self._task_req[tid] for tid in task_ids]
|
||||
req_pool_indices = torch.tensor(req_indices, dtype=torch.long, device=device)
|
||||
seq_lens_t = torch.tensor(seq_lens, dtype=torch.long, device=device)
|
||||
|
||||
if start_pos is not None:
|
||||
seq_len = seq_lens[0]
|
||||
out_cache_loc = self._req_pool.req_to_token[
|
||||
req_pool_indices, start_pos:seq_len
|
||||
]
|
||||
else:
|
||||
write_pos = seq_lens_t - 1
|
||||
out_cache_loc = self._req_pool.req_to_token[
|
||||
req_pool_indices, write_pos
|
||||
].unsqueeze(-1)
|
||||
|
||||
return KVCache(
|
||||
k_buffer=self._storage.k_buffer,
|
||||
v_buffer=self._storage.v_buffer,
|
||||
req_to_token=self._req_pool.req_to_token,
|
||||
req_pool_indices=req_pool_indices,
|
||||
seq_lens=seq_lens_t,
|
||||
out_cache_loc=out_cache_loc,
|
||||
max_len=max(seq_lens),
|
||||
)
|
||||
|
||||
# ---- internals ----
|
||||
|
||||
def _alloc_tokens(self, n: int) -> Optional[List[int]]:
|
||||
if self.page_size != 1:
|
||||
raise RuntimeError("_alloc_tokens is for page_size=1 only")
|
||||
slots = []
|
||||
for _ in range(n):
|
||||
p = self._alloc.alloc()
|
||||
if p < 0:
|
||||
for s in slots:
|
||||
self._alloc.free(s)
|
||||
return None
|
||||
slots.append(p)
|
||||
return slots
|
||||
|
||||
def _write_req_to_token(self, task_id: str, prompt_ids: List[int], cached: int):
|
||||
req_idx = self._task_req[task_id]
|
||||
total = len(prompt_ids)
|
||||
|
||||
if self.contiguous:
|
||||
return
|
||||
|
||||
if self.page_size == 1:
|
||||
slots = self._task_slots.get(task_id, [])
|
||||
all_slots = slots[: total - cached]
|
||||
if all_slots:
|
||||
self._req_pool.req_to_token[req_idx, cached:total] = torch.tensor(
|
||||
all_slots, dtype=torch.long, device=self.device
|
||||
)
|
||||
else:
|
||||
pages = self._task_pages.get(task_id, [])
|
||||
for pos in range(cached, total):
|
||||
page_idx = pos // self.page_size
|
||||
page_offset = pos % self.page_size
|
||||
if page_idx < len(pages):
|
||||
token_slot = pages[page_idx] * self.page_size + page_offset
|
||||
self._req_pool.req_to_token[req_idx, pos] = token_slot
|
||||
|
||||
@@ -3,7 +3,7 @@ from typing import List, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.inference.core.cache import KVCache
|
||||
from astrai.inference.core.cache import PagePool
|
||||
from astrai.inference.core.task import Task
|
||||
from astrai.inference.sample import sample
|
||||
from astrai.model.automodel import AutoModel
|
||||
@@ -19,7 +19,7 @@ class Executor:
|
||||
self,
|
||||
model: AutoModel,
|
||||
tokenizer: AutoTokenizer,
|
||||
kv_cache: KVCache,
|
||||
kv_cache: PagePool,
|
||||
device: Optional[str] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
@@ -57,7 +57,9 @@ class Executor:
|
||||
input_ids,
|
||||
input_mask=input_mask,
|
||||
position_ids=position_ids,
|
||||
paged_cache=self.kv_cache.bind_tasks(task_ids, prompt_len, self.device),
|
||||
kv_cache=self.kv_cache.bind_tasks(
|
||||
task_ids, [prompt_len] * batch_sz, self.device, start_pos=start_pos
|
||||
),
|
||||
)
|
||||
|
||||
def execute_decode(
|
||||
@@ -128,11 +130,10 @@ class Executor:
|
||||
outputs = self.model(
|
||||
input_ids.unsqueeze(1),
|
||||
input_mask=input_mask,
|
||||
paged_cache=self.kv_cache.bind_tasks(
|
||||
kv_cache=self.kv_cache.bind_tasks(
|
||||
task_ids,
|
||||
total_len,
|
||||
[t.next_pos + 1 for t in tasks],
|
||||
self.device,
|
||||
write_positions=position_ids,
|
||||
),
|
||||
position_ids=position_ids.unsqueeze(1),
|
||||
)
|
||||
|
||||
@@ -5,7 +5,7 @@ from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.inference.core.cache import ContiguousCache, KVCache
|
||||
from astrai.inference.core.cache import PagePool
|
||||
from astrai.inference.core.executor import Executor
|
||||
from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
|
||||
from astrai.model.automodel import AutoModel
|
||||
@@ -23,10 +23,9 @@ class InferenceScheduler:
|
||||
tokenizer: AutoTokenizer,
|
||||
max_batch_size: int = 16,
|
||||
max_seq_len: Optional[int] = None,
|
||||
max_prompt_len: int = 2048,
|
||||
device: Optional[str] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
cache: Optional[KVCache] = None,
|
||||
cache: Optional[PagePool] = None,
|
||||
):
|
||||
config = model.config
|
||||
|
||||
@@ -47,21 +46,20 @@ class InferenceScheduler:
|
||||
if cache is not None:
|
||||
self._cache = cache
|
||||
else:
|
||||
self._cache = ContiguousCache(
|
||||
config.num_hidden_layers,
|
||||
max_batch_size,
|
||||
self.max_seq_len,
|
||||
config.num_key_value_heads,
|
||||
head_dim,
|
||||
self.device,
|
||||
self.dtype,
|
||||
self._cache = PagePool(
|
||||
n_layers=config.num_hidden_layers,
|
||||
n_kv_heads=config.num_key_value_heads,
|
||||
head_dim=head_dim,
|
||||
max_batch_size=max_batch_size,
|
||||
max_seq_len=self.max_seq_len,
|
||||
device=self.device,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
|
||||
self._task_mgr = TaskManager(
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=max_batch_size,
|
||||
max_seq_len=self.max_seq_len,
|
||||
max_prompt_len=max_prompt_len,
|
||||
)
|
||||
|
||||
self._executor = Executor(
|
||||
|
||||
@@ -6,6 +6,8 @@ from collections import deque
|
||||
from enum import Enum
|
||||
from typing import Any, Callable, Deque, Dict, List, Optional
|
||||
|
||||
from tokenizers.decoders import DecodeStream
|
||||
|
||||
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -14,37 +16,30 @@ STOP = object()
|
||||
|
||||
|
||||
class StreamDecoder:
|
||||
"""Incremental decoder for byte-level BPE streaming.
|
||||
"""Incremental decoder backed by the tokenizers library's DecodeStream.
|
||||
|
||||
Byte-level BPE may split a single Unicode character (e.g. em-dash,
|
||||
smart quotes) across multiple tokens. Decoding such a token in
|
||||
isolation produces U+FFFD (replacement char). This decoder
|
||||
accumulates token IDs and only emits text once the trailing
|
||||
characters are complete, buffering incomplete multi-byte sequences
|
||||
until the next token arrives.
|
||||
Delegates to the Rust-native streaming decoder which maintains an
|
||||
O(1) bounded token buffer internally (via prefix drain), avoiding
|
||||
the O(n²) cost of re-decoding the full history on each step.
|
||||
|
||||
Multi-byte UTF-8 sequences split across token boundaries are
|
||||
buffered until complete; ``push`` returns "" while the trailing
|
||||
sequence is still incomplete.
|
||||
"""
|
||||
|
||||
__slots__ = ("_tokenizer", "_ids", "_emitted")
|
||||
__slots__ = ("_stream", "_tok")
|
||||
|
||||
def __init__(self, tokenizer: AutoTokenizer):
|
||||
self._tokenizer = tokenizer
|
||||
self._ids: List[int] = []
|
||||
self._emitted: str = ""
|
||||
self._tok = tokenizer._tokenizer
|
||||
self._stream = DecodeStream(skip_special_tokens=True)
|
||||
|
||||
def push(self, token_id: int) -> str:
|
||||
"""Append a token ID and return newly completed text.
|
||||
|
||||
Returns "" while a multi-byte character is still incomplete.
|
||||
"""
|
||||
self._ids.append(token_id)
|
||||
full = self._tokenizer.decode(self._ids, skip_special_tokens=True)
|
||||
if full.endswith("\ufffd"):
|
||||
return ""
|
||||
if len(full) > len(self._emitted):
|
||||
diff = full[len(self._emitted) :]
|
||||
self._emitted = full
|
||||
return diff
|
||||
return ""
|
||||
chunk = self._stream.step(self._tok, token_id)
|
||||
return chunk or ""
|
||||
|
||||
|
||||
class TaskStatus(Enum):
|
||||
@@ -101,18 +96,11 @@ class Task:
|
||||
def flush_remaining(self, tokenizer: AutoTokenizer) -> str:
|
||||
"""Emit any text still buffered in the decoder.
|
||||
|
||||
Called when generation terminates (max_tokens reached, stop
|
||||
sequence, or external removal) to avoid dropping a final
|
||||
incomplete-looking fragment that is actually complete when
|
||||
adjacent to the stop token.
|
||||
With the Rust-native DecodeStream, the stream is always in a
|
||||
correct state — any completed text was already emitted by the
|
||||
last ``push``. A trailing incomplete multi-byte sequence has no
|
||||
valid text to emit, so this is a no-op.
|
||||
"""
|
||||
if self._decoder is None or not self.output_ids:
|
||||
return ""
|
||||
full = tokenizer.decode(self.output_ids, skip_special_tokens=True)
|
||||
if len(full) > len(self._decoder._emitted):
|
||||
diff = full[len(self._decoder._emitted) :]
|
||||
self._decoder._emitted = full
|
||||
return diff
|
||||
return ""
|
||||
|
||||
@property
|
||||
@@ -135,12 +123,10 @@ class TaskManager:
|
||||
tokenizer: AutoTokenizer,
|
||||
max_batch_size: int = 16,
|
||||
max_seq_len: int = 8192,
|
||||
max_prompt_len: int = 512,
|
||||
):
|
||||
self.tokenizer = tokenizer
|
||||
self.max_batch_size = max_batch_size
|
||||
self.max_seq_len = max_seq_len
|
||||
self.max_prompt_len = max_prompt_len
|
||||
|
||||
self.waiting_queue: Deque[Task] = deque()
|
||||
self.active_tasks: List[Task] = []
|
||||
@@ -165,10 +151,10 @@ class TaskManager:
|
||||
) -> str:
|
||||
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
|
||||
prompt_ids = self.tokenizer.encode(prompt)
|
||||
if len(prompt_ids) > self.max_prompt_len:
|
||||
prompt_ids = prompt_ids[-self.max_prompt_len :]
|
||||
if len(prompt_ids) > self.max_seq_len:
|
||||
prompt_ids = prompt_ids[-self.max_seq_len :]
|
||||
|
||||
if len(prompt_ids) >= self.max_seq_len:
|
||||
if len(prompt_ids) > self.max_seq_len:
|
||||
if stream_callback:
|
||||
stream_callback(STOP)
|
||||
return task_id
|
||||
|
||||
@@ -8,7 +8,7 @@ from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Tuple,
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from astrai.inference.core.cache import KVCache
|
||||
from astrai.inference.core.cache import PagePool
|
||||
from astrai.inference.core.scheduler import InferenceScheduler
|
||||
from astrai.inference.core.task import STOP
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
@@ -111,9 +111,7 @@ class InferenceEngine:
|
||||
tokenizer: AutoTokenizer,
|
||||
max_batch_size: int = 1,
|
||||
max_seq_len: Optional[int] = None,
|
||||
max_prompt_len: int = 2048,
|
||||
page_size: int = 128,
|
||||
cache: Optional[KVCache] = None,
|
||||
cache: Optional[PagePool] = None,
|
||||
):
|
||||
self.model = model
|
||||
self.tokenizer = tokenizer
|
||||
@@ -122,7 +120,6 @@ class InferenceEngine:
|
||||
tokenizer=self.tokenizer,
|
||||
max_batch_size=max_batch_size,
|
||||
max_seq_len=max_seq_len,
|
||||
max_prompt_len=max_prompt_len,
|
||||
cache=cache,
|
||||
)
|
||||
|
||||
|
||||
@@ -40,11 +40,12 @@ def _disable_random_init(enable: bool = True):
|
||||
setattr(nn.init, n, fn)
|
||||
|
||||
|
||||
class AutoModel(BaseFactory["AutoModel"], nn.Module):
|
||||
"""
|
||||
Autoregressive language model base class.
|
||||
Provides model loading/saving, registration, and generation.
|
||||
"""
|
||||
class ModelFactory(BaseFactory[nn.Module]):
|
||||
"""Pure factory for model dispatch, separated from nn.Module state."""
|
||||
|
||||
|
||||
class AutoModel(nn.Module):
|
||||
"""Model base class with loading/saving and generation."""
|
||||
|
||||
def __init__(self, config: BaseModelConfig):
|
||||
super().__init__()
|
||||
@@ -68,7 +69,7 @@ class AutoModel(BaseFactory["AutoModel"], nn.Module):
|
||||
config = ConfigFactory.load(raw)
|
||||
model_type = config.model_type or "autoregressive_lm"
|
||||
|
||||
actual_cls = AutoModel.get_component_class(model_type)
|
||||
actual_cls = ModelFactory.get_component_class(model_type)
|
||||
|
||||
with _disable_random_init(enable=disable_random_init):
|
||||
model = actual_cls(config)
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from astrai.model.components.attention import GQA, MLA, repeat_kv
|
||||
from astrai.model.components.attention import GQA, MLA
|
||||
from astrai.model.components.decoder_block import DecoderBlock
|
||||
from astrai.model.components.embedding import Embedding
|
||||
from astrai.model.components.linear import Linear
|
||||
@@ -21,5 +21,4 @@ __all__ = [
|
||||
"RotaryEmbedding",
|
||||
"apply_rotary_emb",
|
||||
"get_rotary_emb",
|
||||
"repeat_kv",
|
||||
]
|
||||
|
||||
@@ -5,24 +5,14 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.extension import attention
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.inference.core.cache import CacheView
|
||||
from astrai.inference.core.cache import KVCache
|
||||
from astrai.model.components.linear import Linear
|
||||
from astrai.model.components.norm import RMSNorm
|
||||
from astrai.model.components.rope import apply_rotary_emb
|
||||
|
||||
|
||||
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
|
||||
bs, slen, n_heads, head_dim = x.shape
|
||||
if n_rep == 1:
|
||||
return x
|
||||
return (
|
||||
x[:, :, :, None, :]
|
||||
.expand(bs, slen, n_heads, n_rep, head_dim)
|
||||
.reshape(bs, slen, n_heads * n_rep, head_dim)
|
||||
)
|
||||
|
||||
|
||||
class AttnFactory(BaseFactory[nn.Module]):
|
||||
pass
|
||||
|
||||
@@ -75,7 +65,7 @@ class GQA(nn.Module):
|
||||
x: Tensor,
|
||||
rotary_emb: Tensor,
|
||||
attn_mask: Tensor = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
kv_cache: Optional[KVCache] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
q = self._split_heads(self.q_proj(x), self.n_heads)
|
||||
@@ -86,19 +76,7 @@ class GQA(nn.Module):
|
||||
if self.use_qk_norm:
|
||||
q, k = self.q_norm(q), self.k_norm(k)
|
||||
|
||||
if paged_cache is not None:
|
||||
paged_cache.write(self.layer_id, k, v)
|
||||
k, v = paged_cache.gather(self.layer_id)
|
||||
|
||||
k, v = repeat_kv(k, self.n_rep), repeat_kv(v, self.n_rep)
|
||||
|
||||
q, k, v = q.permute(0, 2, 1, 3), k.permute(0, 2, 1, 3), v.permute(0, 2, 1, 3)
|
||||
sdqa_out = (
|
||||
F.scaled_dot_product_attention(q, k, v, attn_mask, is_causal=is_causal)
|
||||
.permute(0, 2, 1, 3)
|
||||
.contiguous()
|
||||
.flatten(2)
|
||||
)
|
||||
sdqa_out = attention(q, k, v, kv_cache, self.layer_id, attn_mask, is_causal)
|
||||
|
||||
if self.use_gated_attention:
|
||||
sdqa_out = sdqa_out * F.sigmoid(self.gate(x))
|
||||
@@ -161,7 +139,7 @@ class MLA(nn.Module):
|
||||
x: Tensor,
|
||||
rotary_emb: Tensor,
|
||||
attn_mask: Tensor = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
kv_cache: Optional[KVCache] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
bsz, seq_len, _ = x.size()
|
||||
@@ -193,18 +171,7 @@ class MLA(nn.Module):
|
||||
q = self.q_norm(q)
|
||||
k = self.k_norm(k)
|
||||
|
||||
if paged_cache is not None:
|
||||
paged_cache.write(self.layer_id, k, v)
|
||||
k, v = paged_cache.gather(self.layer_id)
|
||||
|
||||
q = q.permute(0, 2, 1, 3)
|
||||
k = k.permute(0, 2, 1, 3)
|
||||
v = v.permute(0, 2, 1, 3)
|
||||
|
||||
attn_out = F.scaled_dot_product_attention(
|
||||
q, k, v, attn_mask, is_causal=is_causal
|
||||
)
|
||||
attn_out = attn_out.permute(0, 2, 1, 3).contiguous().flatten(2)
|
||||
attn_out = attention(q, k, v, kv_cache, self.layer_id, attn_mask, is_causal)
|
||||
|
||||
if self.use_gated_attention:
|
||||
attn_out = attn_out * F.sigmoid(self.gate(x))
|
||||
|
||||
@@ -4,7 +4,7 @@ from typing import Optional
|
||||
import torch.nn as nn
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.inference.core.cache import CacheView
|
||||
from astrai.inference.core.cache import KVCache
|
||||
from astrai.model.components.attention import AttnFactory
|
||||
from astrai.model.components.mlp import FFNFactory
|
||||
from astrai.model.components.norm import RMSNorm
|
||||
@@ -33,14 +33,14 @@ class DecoderBlock(nn.Module):
|
||||
x: Tensor,
|
||||
rotary_emb: Tensor,
|
||||
attention_mask: Optional[Tensor] = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
kv_cache: Optional[KVCache] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
attn_output = self.attention(
|
||||
self.input_norm(x),
|
||||
rotary_emb,
|
||||
attention_mask,
|
||||
paged_cache,
|
||||
kv_cache,
|
||||
is_causal,
|
||||
)
|
||||
x = attn_output + x
|
||||
|
||||
@@ -1,11 +1,12 @@
|
||||
import logging
|
||||
from dataclasses import asdict, dataclass
|
||||
from dataclasses import asdict
|
||||
from pathlib import Path
|
||||
from typing import Optional, Set
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from pydantic.dataclasses import dataclass
|
||||
|
||||
from astrai.model.components.linear import Linear
|
||||
from astrai.serialization import (
|
||||
|
||||
@@ -5,7 +5,7 @@ import torch.nn as nn
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.config.model_config import EncoderConfig
|
||||
from astrai.model.automodel import AutoModel
|
||||
from astrai.model.automodel import AutoModel, ModelFactory
|
||||
from astrai.model.components.decoder_block import DecoderBlock
|
||||
from astrai.model.components.embedding import Embedding
|
||||
from astrai.model.components.norm import RMSNorm
|
||||
@@ -13,7 +13,7 @@ from astrai.model.components.rope import RotaryEmbedding
|
||||
from astrai.model.transformer import process_attention_mask
|
||||
|
||||
|
||||
@AutoModel.register("embedding")
|
||||
@ModelFactory.register("embedding")
|
||||
class EmbeddingEncoder(AutoModel):
|
||||
def __init__(self, config: EncoderConfig):
|
||||
super().__init__(config)
|
||||
|
||||
@@ -5,8 +5,8 @@ import torch.nn as nn
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.inference.core.cache import CacheView
|
||||
from astrai.model.automodel import AutoModel
|
||||
from astrai.inference.core.cache import KVCache
|
||||
from astrai.model.automodel import AutoModel, ModelFactory
|
||||
from astrai.model.components.decoder_block import DecoderBlock
|
||||
from astrai.model.components.embedding import Embedding
|
||||
from astrai.model.components.linear import Linear
|
||||
@@ -26,7 +26,7 @@ def process_attention_mask(
|
||||
return input_mask
|
||||
|
||||
|
||||
@AutoModel.register("autoregressive_lm")
|
||||
@ModelFactory.register("autoregressive_lm")
|
||||
class AutoRegressiveLM(AutoModel):
|
||||
"""Autoregressive language model with paged KV cache."""
|
||||
|
||||
@@ -103,7 +103,7 @@ class AutoRegressiveLM(AutoModel):
|
||||
self,
|
||||
input_ids: Tensor,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
kv_cache: Optional[KVCache] = None,
|
||||
position_ids: Optional[Tensor] = None,
|
||||
) -> Dict[str, Tensor]:
|
||||
assert input_ids.ndim == 2
|
||||
@@ -114,7 +114,7 @@ class AutoRegressiveLM(AutoModel):
|
||||
use_sdpa_causal_mask = attn_mask is None
|
||||
|
||||
for layer in self.layers:
|
||||
x = layer(x, rotary_emb, attn_mask, paged_cache, use_sdpa_causal_mask)
|
||||
x = layer(x, rotary_emb, attn_mask, kv_cache, use_sdpa_causal_mask)
|
||||
|
||||
hidden_states = self.norm(x)
|
||||
logits = self.lm_head(hidden_states)
|
||||
|
||||
@@ -4,12 +4,12 @@ from astrai.parallel.executor import (
|
||||
BaseExecutor,
|
||||
DDPExecutor,
|
||||
ExecutorFactory,
|
||||
FSDP2Executor,
|
||||
FSDPExecutor,
|
||||
GradientState,
|
||||
NoneExecutor,
|
||||
broadcast_state_dict,
|
||||
create_ref_model,
|
||||
)
|
||||
from astrai.parallel.module import ColumnParallelLinear, RowParallelLinear
|
||||
from astrai.parallel.setup import (
|
||||
get_current_device,
|
||||
get_rank,
|
||||
@@ -26,8 +26,6 @@ __all__ = [
|
||||
"only_on_rank",
|
||||
"setup_parallel",
|
||||
"spawn_parallel_fn",
|
||||
"RowParallelLinear",
|
||||
"ColumnParallelLinear",
|
||||
"ExecutorFactory",
|
||||
"BaseExecutor",
|
||||
"GradientState",
|
||||
@@ -36,5 +34,6 @@ __all__ = [
|
||||
"NoneExecutor",
|
||||
"DDPExecutor",
|
||||
"FSDPExecutor",
|
||||
"FSDP2Executor",
|
||||
"create_ref_model",
|
||||
"broadcast_state_dict",
|
||||
]
|
||||
|
||||
+120
-99
@@ -4,18 +4,15 @@ import contextlib
|
||||
import logging
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
from typing import Any, Callable, Optional, Tuple
|
||||
from typing import Any, Callable, Dict, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
from torch.distributed.fsdp import (
|
||||
FSDPModule,
|
||||
FullStateDictConfig,
|
||||
StateDictType,
|
||||
fully_shard,
|
||||
)
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
from torch.distributed.tensor import DTensor
|
||||
from torch.nn.parallel import DistributedDataParallel as DDP
|
||||
from torch.optim import Optimizer
|
||||
@@ -27,6 +24,82 @@ from astrai.parallel.setup import get_rank, get_world_size
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def broadcast_state_dict(
|
||||
state_dict: Optional[Dict[str, torch.Tensor]],
|
||||
src: int = 0,
|
||||
) -> Optional[Dict[str, torch.Tensor]]:
|
||||
"""Broadcast a state_dict from *src* rank to all ranks.
|
||||
|
||||
Tensors stay on their original device (GPU) for the broadcast.
|
||||
All ranks must call this collectively.
|
||||
|
||||
On non-distributed runs, returns *state_dict* unchanged.
|
||||
"""
|
||||
if not dist.is_initialized() or dist.get_world_size() == 1:
|
||||
return state_dict
|
||||
|
||||
rank = dist.get_rank()
|
||||
|
||||
# Broadcast metadata (keys, shapes, dtypes, device) so non-src ranks
|
||||
# can allocate matching empty tensors on the correct device.
|
||||
if rank == src:
|
||||
device = next(iter(state_dict.values())).device
|
||||
metadata = [
|
||||
(k, tuple(v.shape), v.dtype, str(device)) for k, v in state_dict.items()
|
||||
]
|
||||
else:
|
||||
metadata = None
|
||||
metadata_list = [metadata]
|
||||
dist.broadcast_object_list(metadata_list, src=src)
|
||||
metadata = metadata_list[0]
|
||||
|
||||
# Non-src ranks allocate empty tensors with the broadcasted metadata.
|
||||
if rank != src:
|
||||
state_dict = {
|
||||
k: torch.empty(s, dtype=d, device=torch.device(dev))
|
||||
for k, s, d, dev in metadata
|
||||
}
|
||||
|
||||
# Broadcast each tensor in-place.
|
||||
for tensor in state_dict.values():
|
||||
dist.broadcast(tensor, src=src)
|
||||
|
||||
return state_dict
|
||||
|
||||
|
||||
def create_ref_model(
|
||||
model_fn: Callable[[], nn.Module],
|
||||
executor: Optional["BaseExecutor"] = None,
|
||||
model: Optional[nn.Module] = None,
|
||||
state_dict: Optional[Dict[str, torch.Tensor]] = None,
|
||||
device: Optional[str] = None,
|
||||
) -> Optional[nn.Module]:
|
||||
"""Create a frozen reference model from executor or state dict.
|
||||
|
||||
In distributed mode (FSDP), ``unwrap_model`` returns ``None`` on
|
||||
non-rank-0. The state_dict is broadcast from rank-0 to all ranks
|
||||
so every rank gets a complete copy.
|
||||
"""
|
||||
if state_dict is None and executor is not None and model is not None:
|
||||
state_dict = executor.unwrap_model(model)
|
||||
|
||||
# FSDP's unwrap_model returns None on non-rank-0. Broadcast from
|
||||
# rank-0 so every rank receives a complete state_dict.
|
||||
if executor is not None and executor.use_distributed:
|
||||
state_dict = broadcast_state_dict(state_dict)
|
||||
|
||||
if state_dict is None:
|
||||
return None
|
||||
|
||||
ref_model = model_fn()
|
||||
ref_model.load_state_dict(state_dict)
|
||||
ref_model.requires_grad_(False)
|
||||
ref_model.eval()
|
||||
if device is not None:
|
||||
ref_model = ref_model.to(device=device)
|
||||
return ref_model
|
||||
|
||||
|
||||
class GradientState:
|
||||
def __init__(self, grad_accum_steps: int = 1):
|
||||
self.num_steps = max(grad_accum_steps, 1)
|
||||
@@ -95,11 +168,14 @@ class BaseExecutor:
|
||||
optimizer_fn: Optional[Callable[[nn.Module], Optimizer]] = None,
|
||||
scheduler_fn: Optional[Callable[[Optimizer], LRScheduler]] = None,
|
||||
before_wrap: Optional[Callable[[nn.Module], nn.Module]] = None,
|
||||
after_wrap: Optional[Callable[[nn.Module], nn.Module]] = None,
|
||||
) -> Tuple[nn.Module, Optional[Optimizer], Optional[LRScheduler]]:
|
||||
model = model_fn()
|
||||
if before_wrap is not None:
|
||||
model = before_wrap(model)
|
||||
model = self._prepare_model(model)
|
||||
if after_wrap is not None:
|
||||
model = after_wrap(model)
|
||||
optimizer = None
|
||||
scheduler = None
|
||||
if optimizer_fn is not None:
|
||||
@@ -238,88 +314,11 @@ class DDPExecutor(BaseExecutor):
|
||||
|
||||
@ExecutorFactory.register("fsdp")
|
||||
class FSDPExecutor(BaseExecutor):
|
||||
def __init__(
|
||||
self,
|
||||
grad_accum_steps: int = 1,
|
||||
process_group=None,
|
||||
sharding_strategy=None,
|
||||
cpu_offload=None,
|
||||
auto_wrap_policy=None,
|
||||
backward_prefetch=None,
|
||||
mixed_precision=None,
|
||||
ignored_modules=None,
|
||||
param_init_fn=None,
|
||||
sync_module_states: bool = False,
|
||||
forward_prefetch: bool = False,
|
||||
limit_all_gathers: bool = True,
|
||||
ignored_states=None,
|
||||
device_mesh=None,
|
||||
):
|
||||
super().__init__(grad_accum_steps=grad_accum_steps)
|
||||
self._fsdp_kwargs = {
|
||||
k: v
|
||||
for k, v in dict(
|
||||
process_group=process_group,
|
||||
sharding_strategy=sharding_strategy,
|
||||
cpu_offload=cpu_offload,
|
||||
auto_wrap_policy=auto_wrap_policy,
|
||||
backward_prefetch=backward_prefetch,
|
||||
mixed_precision=mixed_precision,
|
||||
ignored_modules=ignored_modules,
|
||||
param_init_fn=param_init_fn,
|
||||
sync_module_states=sync_module_states,
|
||||
forward_prefetch=forward_prefetch,
|
||||
limit_all_gathers=limit_all_gathers,
|
||||
use_orig_params=True,
|
||||
ignored_states=ignored_states,
|
||||
device_mesh=device_mesh,
|
||||
).items()
|
||||
if v is not None
|
||||
}
|
||||
self._original_model: Optional[nn.Module] = None
|
||||
|
||||
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
||||
if not self.use_distributed:
|
||||
logger.warning("FSDP backend selected but world_size=1, model not wrapped")
|
||||
return model
|
||||
self._original_model = model
|
||||
device_id = torch.device("cuda", get_rank())
|
||||
model = FSDP(model, device_id=device_id, **self._fsdp_kwargs)
|
||||
logger.info("Model wrapped with FSDP (world_size=%d)", get_world_size())
|
||||
return model
|
||||
|
||||
def _no_sync(self, model: nn.Module):
|
||||
if isinstance(model, FSDP):
|
||||
return model.no_sync()
|
||||
return contextlib.nullcontext()
|
||||
|
||||
def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
|
||||
if isinstance(model, FSDP) and self.use_distributed:
|
||||
total_norm = model.clip_grad_norm_(max_norm)
|
||||
if isinstance(total_norm, torch.Tensor):
|
||||
return total_norm.item()
|
||||
return total_norm
|
||||
return super().clip_grad_norm(model, max_norm)
|
||||
|
||||
def unwrap_model(self, model: nn.Module):
|
||||
if isinstance(model, FSDP) and self.use_distributed:
|
||||
with FSDP.state_dict_type(
|
||||
model,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
):
|
||||
return model.state_dict()
|
||||
|
||||
return model.state_dict()
|
||||
|
||||
|
||||
@ExecutorFactory.register("fsdp2")
|
||||
class FSDP2Executor(BaseExecutor):
|
||||
"""FSDP2 executor using `torch.distributed.fsdp.fully_shard` (per-module API).
|
||||
"""FSDP executor using `torch.distributed.fsdp.fully_shard` (per-module API).
|
||||
|
||||
Wraps each child module individually via ``fully_shard``.
|
||||
Skips the root model because ``ABC + Generic[T]`` in the MRO makes
|
||||
FSDP2's dynamic ``__class__`` assignment fail at the CPython level.
|
||||
``fully_shard``'s dynamic ``__class__`` assignment fail at the CPython level.
|
||||
Original ``Parameter`` objects are preserved (as DTensors) — no
|
||||
``FlatParameter``, no ``use_orig_params=True`` hack.
|
||||
"""
|
||||
@@ -329,7 +328,7 @@ class FSDP2Executor(BaseExecutor):
|
||||
grad_accum_steps: int = 1,
|
||||
mesh: Optional[Any] = None,
|
||||
mp_policy: Optional[Any] = None,
|
||||
reshard_after_forward: bool = True,
|
||||
reshard_after_forward: bool = False,
|
||||
):
|
||||
super().__init__(grad_accum_steps=grad_accum_steps)
|
||||
self._mesh = mesh
|
||||
@@ -338,7 +337,7 @@ class FSDP2Executor(BaseExecutor):
|
||||
|
||||
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
||||
if not self.use_distributed:
|
||||
logger.warning("FSDP2 backend selected but world_size=1, model not wrapped")
|
||||
logger.warning("FSDP backend selected but world_size=1, model not wrapped")
|
||||
return model
|
||||
|
||||
kwargs = dict(
|
||||
@@ -356,7 +355,7 @@ class FSDP2Executor(BaseExecutor):
|
||||
fully_shard(child, **kwargs)
|
||||
|
||||
logger.info(
|
||||
"FSDP2 wrapping applied to %d direct children (root skipped for ABC compat)",
|
||||
"FSDP wrapping applied to %d direct children (root skipped for ABC compat)",
|
||||
len(list(model.children())),
|
||||
)
|
||||
return model
|
||||
@@ -376,32 +375,54 @@ class FSDP2Executor(BaseExecutor):
|
||||
yield
|
||||
|
||||
def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
|
||||
if self.use_distributed:
|
||||
total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)
|
||||
if isinstance(total_norm, torch.Tensor):
|
||||
return total_norm.item()
|
||||
return total_norm
|
||||
return super().clip_grad_norm(model, max_norm)
|
||||
if not self.use_distributed:
|
||||
return super().clip_grad_norm(model, max_norm)
|
||||
|
||||
# FSDP params are DTensors (sharded across ranks).
|
||||
# torch.nn.utils.clip_grad_norm_ computes LOCAL norm per rank,
|
||||
# so we must all-reduce to get the global norm before clipping.
|
||||
local_norm = torch.nn.utils.get_total_norm(
|
||||
[p.grad for p in model.parameters() if p.grad is not None],
|
||||
)
|
||||
if isinstance(local_norm, DTensor):
|
||||
local_norm = local_norm.to_local()
|
||||
total_norm_sq = local_norm**2
|
||||
dist.all_reduce(total_norm_sq, op=dist.ReduceOp.SUM)
|
||||
total_norm = total_norm_sq.sqrt()
|
||||
|
||||
clip_coef = max_norm / (total_norm + 1e-6)
|
||||
clip_coef_clamped = torch.clamp(clip_coef, max=1.0)
|
||||
for p in model.parameters():
|
||||
if p.grad is not None:
|
||||
p.grad.mul_(clip_coef_clamped)
|
||||
|
||||
return total_norm.item()
|
||||
|
||||
def unwrap_model(self, model: nn.Module):
|
||||
if not self.use_distributed:
|
||||
return model.state_dict()
|
||||
|
||||
if get_rank() != 0:
|
||||
return None
|
||||
|
||||
# unshard() and full_tensor() are collective ops — all ranks must
|
||||
# participate. Non-rank-0 ranks still call them but discard results.
|
||||
for module in model.modules():
|
||||
if isinstance(module, FSDPModule):
|
||||
module.unshard()
|
||||
|
||||
state_dict = model.state_dict()
|
||||
result = {
|
||||
k: (v.full_tensor() if isinstance(v, DTensor) else v)
|
||||
for k, v in state_dict.items()
|
||||
}
|
||||
result = {}
|
||||
for k, v in state_dict.items():
|
||||
if isinstance(v, DTensor):
|
||||
full = v.full_tensor()
|
||||
if get_rank() == 0:
|
||||
result[k] = full
|
||||
elif get_rank() == 0:
|
||||
result[k] = v
|
||||
|
||||
for module in model.modules():
|
||||
if isinstance(module, FSDPModule):
|
||||
module.reshard()
|
||||
|
||||
if get_rank() != 0:
|
||||
return None
|
||||
|
||||
return result
|
||||
|
||||
@@ -1,115 +0,0 @@
|
||||
from typing import Dict
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
class ParallelModel(nn.Module):
|
||||
def __init__(self, process_group: dist.ProcessGroup):
|
||||
super().__init__()
|
||||
self.process_group = process_group
|
||||
self.rank = dist.get_rank(self.process_group)
|
||||
self.world_size = dist.get_world_size(self.process_group)
|
||||
|
||||
|
||||
class RowParallelLinear(ParallelModel):
|
||||
def __init__(
|
||||
self,
|
||||
process_group: dist.ProcessGroup,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
bias: bool = True,
|
||||
reduce_results: bool = True,
|
||||
):
|
||||
super().__init__(process_group)
|
||||
|
||||
self.in_features = in_features
|
||||
self.out_features = out_features
|
||||
self.in_features_per_rank = in_features // self.world_size
|
||||
self.reduce_results = reduce_results
|
||||
|
||||
if in_features % self.world_size != 0:
|
||||
raise ValueError(
|
||||
f"in_features must be divisible by world_size. Got {in_features} and {self.world_size}"
|
||||
)
|
||||
|
||||
self.weight = nn.Parameter(torch.empty(out_features, self.in_features_per_rank))
|
||||
self.bias = nn.Parameter(torch.zeros(out_features)) if bias else None
|
||||
|
||||
def forward(self, input: Tensor) -> Tensor:
|
||||
output = F.linear(input, self.weight)
|
||||
|
||||
if self.reduce_results:
|
||||
dist.all_reduce(output, op=dist.ReduceOp.SUM, group=self.process_group)
|
||||
|
||||
if self.bias is not None:
|
||||
output += self.bias
|
||||
|
||||
return output
|
||||
|
||||
def load_state_dict(self, state_dict: Dict[str, Tensor]):
|
||||
full_weight = state_dict.get("weight")
|
||||
full_bias = state_dict.get("bias")
|
||||
|
||||
start_idx = self.rank * self.in_features_per_rank
|
||||
end_idx = start_idx + self.in_features_per_rank
|
||||
weight_slice = full_weight[:, start_idx:end_idx]
|
||||
self.weight.data.copy_(weight_slice)
|
||||
|
||||
if self.bias is not None:
|
||||
self.bias.data.copy_(full_bias)
|
||||
|
||||
|
||||
class ColumnParallelLinear(ParallelModel):
|
||||
def __init__(
|
||||
self,
|
||||
process_group: dist.ProcessGroup,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
bias: bool = True,
|
||||
gather_results: bool = True,
|
||||
):
|
||||
super().__init__(process_group)
|
||||
|
||||
self.in_features = in_features
|
||||
self.out_features = out_features
|
||||
self.out_features_per_rank = out_features // self.world_size
|
||||
self.gather_results = gather_results
|
||||
|
||||
if out_features % self.world_size != 0:
|
||||
raise ValueError(
|
||||
f"out_features must be divisible by world_size. Got {out_features} and {self.world_size}"
|
||||
)
|
||||
|
||||
self.weight = nn.Parameter(
|
||||
torch.empty(self.out_features_per_rank, self.in_features)
|
||||
)
|
||||
self.bias = (
|
||||
nn.Parameter(torch.zeros(self.out_features_per_rank)) if bias else None
|
||||
)
|
||||
|
||||
def forward(self, input: Tensor) -> Tensor:
|
||||
output = F.linear(input, self.weight, self.bias)
|
||||
|
||||
if self.gather_results:
|
||||
output_list = [torch.empty_like(output) for _ in range(self.world_size)]
|
||||
dist.all_gather(output_list, output, group=self.process_group)
|
||||
output = torch.cat(output_list, dim=-1)
|
||||
|
||||
return output
|
||||
|
||||
def load_state_dict(self, state_dict: Dict[str, Tensor]):
|
||||
full_weight = state_dict.get("weight")
|
||||
full_bias = state_dict.get("bias")
|
||||
|
||||
start_idx = self.rank * self.out_features_per_rank
|
||||
end_idx = start_idx + self.out_features_per_rank
|
||||
weight_slice = full_weight[start_idx:end_idx, :]
|
||||
self.weight.data.copy_(weight_slice)
|
||||
|
||||
if self.bias is not None:
|
||||
bias_slice = full_bias[start_idx:end_idx]
|
||||
self.bias.data.copy_(bias_slice)
|
||||
@@ -12,7 +12,7 @@ import torch
|
||||
import torch.distributed as dist
|
||||
import torch.multiprocessing as mp
|
||||
|
||||
from astrai.parallel.signal_handler import install_early_signal_handlers
|
||||
from astrai.signal_handler import install_early_signal_handlers
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""Config-driven JSONL preprocessing pipeline.
|
||||
|
||||
Composes a :class:`BaseMaskBuilder` (selected by ``input.type``) with
|
||||
sharding and flush to ``.h5`` / ``.bin`` storage. Packing, position-id
|
||||
sharding and flush to ``.bin`` storage. Packing, position-id
|
||||
generation and storage writing are each delegated to pluggable strategies,
|
||||
dispatched by configuration keys.
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""Storage writer strategies for pipeline output.
|
||||
|
||||
The :class:`StoreWriter` abstraction decouples the pipeline from the
|
||||
concrete storage format (bin / h5). The pipeline builds a ``{key:
|
||||
concrete storage format (bin). The pipeline builds a ``{key:
|
||||
List[Tensor]}`` dict and delegates the write to the writer selected
|
||||
by ``output.storage_format``.
|
||||
"""
|
||||
@@ -15,7 +15,7 @@ from typing import Dict, List
|
||||
import torch
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.serialization import save_bin, save_h5
|
||||
from astrai.serialization import save_bin
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -54,22 +54,3 @@ class BinWriter(StoreWriter):
|
||||
exc_info=True,
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
@StoreWriterFactory.register("h5")
|
||||
class H5Writer(StoreWriter):
|
||||
def save(self, output_dir, domain, shard_idx, tensors):
|
||||
chunk_dir = os.path.join(output_dir, domain)
|
||||
file_path = os.path.join(chunk_dir, f"data_{shard_idx:04d}.h5")
|
||||
try:
|
||||
save_h5(chunk_dir, f"data_{shard_idx:04d}", tensors)
|
||||
except Exception:
|
||||
if os.path.exists(file_path):
|
||||
os.remove(file_path)
|
||||
logger.error(
|
||||
"Failed to write shard %s/data_%04d.h5, cleaned up partial output",
|
||||
domain,
|
||||
shard_idx,
|
||||
exc_info=True,
|
||||
)
|
||||
raise
|
||||
|
||||
@@ -20,9 +20,7 @@ from astrai.serialization.checkpoint import (
|
||||
from astrai.serialization.dataset import (
|
||||
load_bin,
|
||||
load_bin_offsets,
|
||||
load_h5,
|
||||
save_bin,
|
||||
save_h5,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
@@ -39,7 +37,5 @@ __all__ = [
|
||||
"save_torch",
|
||||
"load_bin",
|
||||
"load_bin_offsets",
|
||||
"load_h5",
|
||||
"save_bin",
|
||||
"save_h5",
|
||||
]
|
||||
|
||||
@@ -1,55 +1,14 @@
|
||||
"""Dataset storage serialization helpers (HDF5 / memory-mapped binary)."""
|
||||
"""Dataset storage serialization helpers (memory-mapped binary)."""
|
||||
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import h5py
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
def save_h5(file_path: str, file_name: str, tensor_group: Dict[str, List[Tensor]]):
|
||||
os.makedirs(file_path, exist_ok=True)
|
||||
full_file_path = os.path.join(file_path, f"{file_name}.h5")
|
||||
with h5py.File(full_file_path, "w") as f:
|
||||
for key, tensors in tensor_group.items():
|
||||
grp = f.create_group(key)
|
||||
for idx, tensor in enumerate(tensors):
|
||||
arr = tensor.cpu().numpy()
|
||||
grp.create_dataset(f"data_{idx}", data=arr)
|
||||
|
||||
|
||||
def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
|
||||
tensor_group: Dict[str, List[Tensor]] = {}
|
||||
|
||||
root_path = Path(file_path)
|
||||
if root_path.is_file() and root_path.suffix in (".h5", ".hdf5"):
|
||||
h5_files = [root_path]
|
||||
else:
|
||||
h5_files = list(root_path.rglob("*.h5")) + list(root_path.rglob("*.hdf5"))
|
||||
|
||||
for h5_file in h5_files:
|
||||
with h5py.File(h5_file, "r") as f:
|
||||
for key in f.keys():
|
||||
grp = f[key]
|
||||
dsets = []
|
||||
for dset_name in grp.keys():
|
||||
dset = grp[dset_name]
|
||||
tensor = torch.from_numpy(dset[:])
|
||||
if share_memory:
|
||||
tensor = tensor.share_memory_()
|
||||
dsets.append(tensor)
|
||||
|
||||
if tensor_group.get(key) is None:
|
||||
tensor_group[key] = []
|
||||
tensor_group[key].extend(dsets)
|
||||
|
||||
return tensor_group
|
||||
|
||||
|
||||
def save_bin(
|
||||
file_path: str,
|
||||
tensor_group: Dict[str, List[Tensor]],
|
||||
@@ -65,7 +24,7 @@ def save_bin(
|
||||
offsets, preserving backward compatibility.
|
||||
|
||||
Nested keys (``List[List[Tensor]]`` such as GRPO ``responses``) are
|
||||
not supported in bin format — use H5 for those.
|
||||
not supported in bin format — use JSONL for those.
|
||||
"""
|
||||
os.makedirs(file_path, exist_ok=True)
|
||||
record_keys = set(record_keys or [])
|
||||
@@ -74,7 +33,7 @@ def save_bin(
|
||||
if tensors and isinstance(tensors[0], list):
|
||||
raise ValueError(
|
||||
f"Nested key '{key}' (List[List[Tensor]]) is not supported "
|
||||
f"in bin format. Use H5 or JSONL storage instead."
|
||||
f"in bin format. Use JSONL storage instead."
|
||||
)
|
||||
cat = torch.cat(tensors, dim=0)
|
||||
entry: Dict[str, Any] = {
|
||||
@@ -112,7 +71,7 @@ def load_bin_offsets(file_path: str) -> Dict[str, List[int]]:
|
||||
|
||||
Returns an empty dict when no key has offsets (legacy bin files),
|
||||
in which case record-mode access falls back to per-record segment
|
||||
indexing (H5/JSONL layout).
|
||||
indexing (JSONL layout).
|
||||
"""
|
||||
with open(os.path.join(file_path, "meta.json"), "r") as f:
|
||||
meta = json.load(f)
|
||||
|
||||
@@ -38,12 +38,27 @@ class ChatTemplate:
|
||||
The compiled :class:`~jinja2.Template` holds a dynamically-generated
|
||||
``root`` render function whose ``__module__`` is ``None``; under
|
||||
``pickle`` it falls back to ``__main__`` and breaks ``spawn``-based
|
||||
multiprocessing. By deferring compilation to first access, the
|
||||
default pickle protocol serialises only ``template_str``; each
|
||||
worker rebuilds the cache on first render.
|
||||
multiprocessing. :meth:`__getstate__` drops the cached template so
|
||||
that pickle serialises only ``template_str``; each worker rebuilds
|
||||
the cache on first render.
|
||||
"""
|
||||
return Template(self.template_str)
|
||||
|
||||
def __getstate__(self) -> Dict[str, Any]:
|
||||
"""Exclude the cached Jinja2 template from pickling.
|
||||
|
||||
``Template.root_render_func`` is a dynamically generated closure
|
||||
that cannot be pickled by reference. Dropping ``_compiled`` here
|
||||
lets :class:`cached_property` rebuild it on first access after
|
||||
unpickle.
|
||||
"""
|
||||
state = self.__dict__.copy()
|
||||
state.pop("_compiled", None)
|
||||
return state
|
||||
|
||||
def __setstate__(self, state: Dict[str, Any]) -> None:
|
||||
self.__dict__.update(state)
|
||||
|
||||
@classmethod
|
||||
def from_string(
|
||||
cls,
|
||||
|
||||
@@ -20,8 +20,6 @@ Messages = List[Message]
|
||||
class AutoTokenizer:
|
||||
"""Base tokenizer class with automatic loading support"""
|
||||
|
||||
TOKENIZER_CLASSES = {} # Registry for auto-loading
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path: Optional[Union[str, Path]] = None,
|
||||
@@ -108,17 +106,6 @@ class AutoTokenizer:
|
||||
with open(save_path / "tokenizer_config.json", "w", encoding="utf-8") as f:
|
||||
json.dump(config, f, ensure_ascii=False, indent=2)
|
||||
|
||||
@classmethod
|
||||
def register_tokenizer(cls, name: str, tokenizer_class: type):
|
||||
"""
|
||||
Register a new tokenizer class.
|
||||
|
||||
Args:
|
||||
name: Name to register the tokenizer class under
|
||||
tokenizer_class: The tokenizer class to register
|
||||
"""
|
||||
cls.TOKENIZER_CLASSES[name] = tokenizer_class
|
||||
|
||||
def encode(
|
||||
self,
|
||||
tokens: Union[str, List[str]],
|
||||
|
||||
@@ -9,20 +9,10 @@ import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.parallel.executor import broadcast_state_dict
|
||||
from astrai.trainer.rollout import RolloutResult
|
||||
|
||||
|
||||
def create_ref_model(
|
||||
model_fn: Callable[[], nn.Module], state_dict: Dict[str, Tensor]
|
||||
) -> nn.Module:
|
||||
"""Create a frozen reference model from model_fn + full state dict."""
|
||||
ref_model = model_fn()
|
||||
ref_model.load_state_dict(state_dict)
|
||||
ref_model.requires_grad_(False)
|
||||
ref_model.eval()
|
||||
return ref_model
|
||||
|
||||
|
||||
def move_to_device(batch: Dict[str, Tensor], device: str) -> Dict[str, Tensor]:
|
||||
"""Move batch tensors to specified device with non-blocking transfer."""
|
||||
return {key: value.to(device, non_blocking=True) for key, value in batch.items()}
|
||||
@@ -401,7 +391,11 @@ class GRPOStrategy(BaseStrategy):
|
||||
|
||||
def sync_old_model(self):
|
||||
"""Copy current policy weights to old model."""
|
||||
self.old_model.load_state_dict(self.executor.unwrap_model(self.model))
|
||||
state_dict = self.executor.unwrap_model(self.model)
|
||||
if self.executor.use_distributed:
|
||||
state_dict = broadcast_state_dict(state_dict)
|
||||
if state_dict is not None:
|
||||
self.old_model.load_state_dict(state_dict)
|
||||
|
||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||
batch = move_to_device(batch, self.device)
|
||||
@@ -510,5 +504,5 @@ class GRPOStrategy(BaseStrategy):
|
||||
# Factory aliases: online variants use the same strategy class; the
|
||||
# ``RolloutRunner`` is injected by ``TrainContextBuilder`` to enable
|
||||
# online mode, so no separate subclass is needed.
|
||||
StrategyFactory._entries["online_grpo"] = GRPOStrategy
|
||||
StrategyFactory._entries["online_dpo"] = DPOStrategy
|
||||
StrategyFactory.register("online_grpo")(GRPOStrategy)
|
||||
StrategyFactory.register("online_dpo")(DPOStrategy)
|
||||
|
||||
@@ -235,7 +235,7 @@ class ProgressBarCallback(TrainCallback):
|
||||
class MetricCallback(TrainCallback):
|
||||
def __init__(
|
||||
self,
|
||||
log_dir: str,
|
||||
ckpt_dir: str,
|
||||
save_interval: int,
|
||||
metrics: List[str] = None,
|
||||
val_step: int = 0,
|
||||
@@ -246,8 +246,7 @@ class MetricCallback(TrainCallback):
|
||||
self.val_step = val_step
|
||||
self._next_val_step = 0
|
||||
|
||||
self.log_dir = Path(log_dir) if log_dir else Path.cwd() / "logs"
|
||||
self.log_dir.mkdir(parents=True, exist_ok=True)
|
||||
self.ckpt_dir = Path(ckpt_dir) if ckpt_dir else Path.cwd() / "checkpoint"
|
||||
|
||||
self.log_cache = []
|
||||
|
||||
@@ -306,7 +305,7 @@ class MetricCallback(TrainCallback):
|
||||
|
||||
@only_on_rank(0)
|
||||
def _flush(self, epoch, step):
|
||||
log_file = self.log_dir / f"epoch_{epoch}_step_{step}_metric.jsonl"
|
||||
log_file = self.ckpt_dir / f"epoch_{epoch}_step_{step}" / "metric.jsonl"
|
||||
log_file.parent.mkdir(parents=True, exist_ok=True)
|
||||
with open(log_file, "w") as f:
|
||||
for log in self.log_cache:
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import logging
|
||||
import threading
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
@@ -11,13 +12,15 @@ from astrai.config.train_config import TrainConfig
|
||||
from astrai.dataset import RDSampler
|
||||
from astrai.inference.core.scheduler import InferenceScheduler
|
||||
from astrai.model.components.lora import inject_lora
|
||||
from astrai.parallel.executor import BaseExecutor, ExecutorFactory
|
||||
from astrai.parallel.executor import BaseExecutor, ExecutorFactory, create_ref_model
|
||||
from astrai.parallel.setup import get_current_device, get_rank, get_world_size
|
||||
from astrai.protocols import OptimizerProtocol, SchedulerProtocol
|
||||
from astrai.serialization import Checkpoint, load_json
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
from astrai.trainer.rollout import RolloutGenerator, RolloutRunner
|
||||
from astrai.trainer.strategy import BaseStrategy, StrategyFactory, create_ref_model
|
||||
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -101,18 +104,13 @@ class TrainContextBuilder:
|
||||
if checkpoint.config:
|
||||
model_config = checkpoint.config
|
||||
if self._resume:
|
||||
preloaded_epoch = checkpoint.epoch or cfg.start_epoch
|
||||
if checkpoint.consumed_samples > 0:
|
||||
per_step = (
|
||||
cfg.batch_per_device
|
||||
* get_world_size()
|
||||
* cfg.grad_accum_steps
|
||||
)
|
||||
preloaded_consumed = (
|
||||
checkpoint.consumed_samples // per_step
|
||||
) * per_step
|
||||
else:
|
||||
preloaded_consumed = cfg.start_samples * get_world_size()
|
||||
preloaded_epoch = checkpoint.epoch
|
||||
per_step = (
|
||||
cfg.batch_per_device * get_world_size() * cfg.grad_accum_steps
|
||||
)
|
||||
preloaded_consumed = (
|
||||
checkpoint.consumed_samples // per_step
|
||||
) * per_step
|
||||
preloaded_checkpoint = checkpoint
|
||||
|
||||
if not model_config and hasattr(cfg.model_fn(), "config"):
|
||||
@@ -131,6 +129,12 @@ class TrainContextBuilder:
|
||||
m.load_state_dict(preloaded_state_dict, strict=False)
|
||||
return m
|
||||
|
||||
def _after_wrap(m):
|
||||
if cfg.compile_mode is not None:
|
||||
logger.info("torch.compile enabled (mode=%s)", cfg.compile_mode)
|
||||
m = torch.compile(m, mode=cfg.compile_mode)
|
||||
return m
|
||||
|
||||
context = TrainContext(
|
||||
world_size=get_world_size(),
|
||||
rank=get_rank(),
|
||||
@@ -147,6 +151,7 @@ class TrainContextBuilder:
|
||||
cfg.optimizer_fn,
|
||||
cfg.scheduler_fn,
|
||||
before_wrap=_before_wrap,
|
||||
after_wrap=_after_wrap,
|
||||
)
|
||||
|
||||
train_dataset = cfg.dataset
|
||||
@@ -162,6 +167,15 @@ class TrainContextBuilder:
|
||||
)
|
||||
|
||||
sampler_offset = context.consumed_samples // context.world_size
|
||||
|
||||
if self._resume and sampler_offset > 0:
|
||||
offset = context.world_size - 1
|
||||
num_samples_per_replica = (
|
||||
len(train_dataset) + offset
|
||||
) // context.world_size
|
||||
if num_samples_per_replica > 0:
|
||||
context.epoch = sampler_offset // num_samples_per_replica
|
||||
|
||||
sampler = RDSampler(
|
||||
data_source=train_dataset,
|
||||
start_epoch=context.epoch,
|
||||
@@ -215,17 +229,14 @@ class TrainContextBuilder:
|
||||
needs_old = cfg.strategy in ("grpo", "online_grpo")
|
||||
|
||||
if needs_ref:
|
||||
ref_model = create_ref_model(
|
||||
cfg.model_fn, executor.unwrap_model(context.model)
|
||||
).to(device=device)
|
||||
strategy_kwargs["ref_model"] = ref_model
|
||||
strategy_kwargs["ref_model"] = create_ref_model(
|
||||
cfg.model_fn, executor=executor, model=context.model, device=device
|
||||
)
|
||||
|
||||
old_model = None
|
||||
if needs_old:
|
||||
old_model = create_ref_model(
|
||||
cfg.model_fn, executor.unwrap_model(context.model)
|
||||
).to(device=device)
|
||||
strategy_kwargs["old_model"] = old_model
|
||||
strategy_kwargs["old_model"] = create_ref_model(
|
||||
cfg.model_fn, executor=executor, model=context.model, device=device
|
||||
)
|
||||
|
||||
context.strategy = StrategyFactory.create(
|
||||
cfg.strategy,
|
||||
@@ -257,7 +268,6 @@ class TrainContextBuilder:
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=rollout_batch_size,
|
||||
max_seq_len=max_seq_len,
|
||||
max_prompt_len=max_seq_len or 4096,
|
||||
)
|
||||
|
||||
generator = RolloutGenerator(
|
||||
|
||||
@@ -5,7 +5,7 @@ import torch.distributed as dist
|
||||
|
||||
from astrai.config import TrainConfig
|
||||
from astrai.parallel.setup import spawn_parallel_fn
|
||||
from astrai.parallel.signal_handler import (
|
||||
from astrai.signal_handler import (
|
||||
register_signal_handlers,
|
||||
unregister_signal_handlers,
|
||||
)
|
||||
@@ -42,7 +42,7 @@ class Trainer:
|
||||
),
|
||||
CallbackFactory.create(
|
||||
"metric",
|
||||
log_dir=cfg.log_dir,
|
||||
ckpt_dir=cfg.ckpt_dir,
|
||||
save_interval=cfg.ckpt_interval,
|
||||
metrics=cfg.metrics,
|
||||
val_step=cfg.val_step,
|
||||
|
||||
@@ -19,9 +19,11 @@ struct AttentionParams {
|
||||
// KV strides (K and V share the same layout — only base pointers differ)
|
||||
int kv_stride_b, kv_stride_h, kv_stride_l, kv_stride_d;
|
||||
|
||||
// Mask: 2D [batch, kv_len] (mask_q_stride=0) or 3D [batch, q_len, kv_len]
|
||||
int mask_b_stride; // = kv_len (both 2D and 3D)
|
||||
int mask_q_stride; // 2D: 0 (all q rows share); 3D: kv_len
|
||||
// Mask: 2D [batch, kv_len], 3D [batch, q_len, kv_len],
|
||||
// or 4D [batch, n_heads, q_len, kv_len] (head dim broadcasts when stride=0)
|
||||
int mask_b_stride; // batch stride
|
||||
int mask_h_stride; // head stride (0 = broadcast across heads)
|
||||
int mask_q_stride; // q stride (0 = all q rows share)
|
||||
|
||||
const T* __restrict__ q;
|
||||
const T* __restrict__ k;
|
||||
@@ -52,8 +54,9 @@ struct PagedAttentionParams {
|
||||
// Q strides (layout-agnostic)
|
||||
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
|
||||
|
||||
// Mask strides (2D or 3D)
|
||||
// Mask strides (2D, 3D, or 4D)
|
||||
int mask_b_stride;
|
||||
int mask_h_stride;
|
||||
int mask_q_stride;
|
||||
|
||||
const T* __restrict__ q;
|
||||
|
||||
@@ -24,7 +24,7 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
||||
|
||||
// KV: [batch, kv_head, kv_len, head_dim] — stride-based base
|
||||
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||
int mask_base = batch * p.mask_b_stride;
|
||||
int mask_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
|
||||
|
||||
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
|
||||
|
||||
@@ -70,8 +70,8 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
||||
}
|
||||
|
||||
float new_m = fmaxf(m, partial);
|
||||
float alpha = expf(m - new_m);
|
||||
float beta = expf(partial - new_m);
|
||||
float alpha = __expf(m - new_m);
|
||||
float beta = __expf(partial - new_m);
|
||||
d = d * alpha + beta;
|
||||
|
||||
int v_off = kv_base + kv_idx * p.kv_stride_l
|
||||
@@ -116,8 +116,8 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
|
||||
if (mi <= -FLT_MAX) continue;
|
||||
float li = mlp[s * 2 + 1];
|
||||
float nm = fmaxf(m, mi);
|
||||
float corr = expf(m - nm);
|
||||
float e = expf(mi - nm);
|
||||
float corr = __expf(m - nm);
|
||||
float e = __expf(mi - nm);
|
||||
acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
|
||||
l = fmaf(l, corr, li * e);
|
||||
m = nm;
|
||||
|
||||
@@ -109,8 +109,8 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
||||
int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
||||
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
|
||||
0, 0,
|
||||
p.mask_b_stride, 0,
|
||||
batch,
|
||||
p.mask_b_stride, 0, 0,
|
||||
batch, 0,
|
||||
p.mask,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
|
||||
|
||||
@@ -15,11 +15,15 @@
|
||||
#endif
|
||||
|
||||
// Split-KV: compute number of splits to fill all SMs for small-batch decode.
|
||||
inline int compute_num_splits(int base_blocks, int tiles_total) {
|
||||
// Caps splits so each split processes at least `min_tiles_per_split` tiles,
|
||||
// avoiding excessive loop/prologue overhead when tiles are small.
|
||||
inline int compute_num_splits(int base_blocks, int tiles_total,
|
||||
int min_tiles_per_split = 1) {
|
||||
int sm_count = 0;
|
||||
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
|
||||
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
|
||||
return std::max(1, std::min(n, std::min(tiles_total, MAX_SPLITS)));
|
||||
int max_by_work = tiles_total / min_tiles_per_split;
|
||||
return std::max(1, std::min(n, std::min(max_by_work, MAX_SPLITS)));
|
||||
}
|
||||
|
||||
// ======================================================================
|
||||
@@ -75,15 +79,20 @@ static inline void dispatch_prefill(AttentionParams<bf16>& p) {
|
||||
// ======================================================================
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
// BC=16: halves smem (16KB vs 32KB) → doubles occupancy (6 vs 3 blocks/SM).
|
||||
// For D=256, BC=16 also reduces register pressure (fewer Sacc/PV frags),
|
||||
// enabling STAGES=2 (double-buffer) within the 32KB smem budget — eliminates
|
||||
// the 176-byte spill that STAGES=1+BC=32 suffered.
|
||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||
static inline void launch_decode_mma(AttentionParams<bf16>& p, int group_size) {
|
||||
int G = p.q_head / p.kv_head;
|
||||
constexpr int MAX_G = 16;
|
||||
int num_passes = (G + MAX_G - 1) / MAX_G;
|
||||
int tiles_total = (p.kv_len + 32 - 1) / 32;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
||||
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
|
||||
using Traits = KernelTraits<HEAD_DIM, 32, 1, STAGES>;
|
||||
constexpr int BC = 16;
|
||||
int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total, 2);
|
||||
constexpr int STAGES = 2;
|
||||
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
|
||||
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
|
||||
attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask><<<grid, 32>>>(p);
|
||||
}
|
||||
@@ -136,23 +145,14 @@ template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||
static inline void launch_paged_decode_mma(PagedAttentionParams<bf16>& p, int group_size) {
|
||||
int G = p.q_head / p.kv_head;
|
||||
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;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
||||
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
|
||||
using Traits = KernelTraits<HEAD_DIM, 32, 1, STAGES>;
|
||||
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);
|
||||
} else {
|
||||
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
|
||||
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
|
||||
dim3 block(32, group_size);
|
||||
paged_attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
|
||||
}
|
||||
constexpr int BC = 16;
|
||||
int num_passes = (G + MAX_G - 1) / MAX_G;
|
||||
int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total, 2);
|
||||
constexpr int STAGES = 2;
|
||||
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
|
||||
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);
|
||||
}
|
||||
#endif
|
||||
|
||||
|
||||
@@ -44,6 +44,9 @@ inline void extract_q_dims_and_strides(torch::Tensor& q, int64_t layout, P& p) {
|
||||
}
|
||||
|
||||
// ---- Shared mask packing ----
|
||||
// Accepts 2D [batch, kv_len], 3D [batch, q_len, kv_len],
|
||||
// or 4D [batch, n_heads, q_len, kv_len].
|
||||
// Head/q dimensions with size 1 broadcast (stride set to 0).
|
||||
template <typename P>
|
||||
inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
|
||||
if (p.use_mask) {
|
||||
@@ -54,18 +57,26 @@ inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
|
||||
TORCH_CHECK(m.size(m.dim() - 1) == p.kv_len, "mask kv_len mismatch");
|
||||
if (m.dim() == 2) {
|
||||
p.mask_b_stride = (int)m.stride(0);
|
||||
p.mask_h_stride = 0;
|
||||
p.mask_q_stride = 0;
|
||||
} else if (m.dim() == 3) {
|
||||
TORCH_CHECK(m.size(1) == p.q_len, "mask q_len mismatch");
|
||||
TORCH_CHECK(m.size(1) == 1 || m.size(1) == p.q_len, "mask q_len mismatch");
|
||||
p.mask_b_stride = (int)m.stride(0);
|
||||
p.mask_q_stride = (int)m.stride(1);
|
||||
p.mask_h_stride = 0;
|
||||
p.mask_q_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
|
||||
} else if (m.dim() == 4) {
|
||||
TORCH_CHECK(m.size(2) == 1 || m.size(2) == p.q_len, "mask q_len mismatch");
|
||||
p.mask_b_stride = (int)m.stride(0);
|
||||
p.mask_h_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
|
||||
p.mask_q_stride = (m.size(2) == 1) ? 0 : (int)m.stride(2);
|
||||
} else {
|
||||
TORCH_CHECK(false, "mask must be 2D [batch, kv_len] or 3D [batch, q_len, kv_len]");
|
||||
TORCH_CHECK(false, "mask must be 2D, 3D, or 4D");
|
||||
}
|
||||
p.mask = m.data_ptr<bool>();
|
||||
} else {
|
||||
p.mask = nullptr;
|
||||
p.mask_b_stride = 0;
|
||||
p.mask_h_stride = 0;
|
||||
p.mask_q_stride = 0;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -192,8 +192,8 @@ __device__ inline void mma_softmax_tile(
|
||||
int kv0,
|
||||
int maxc0, int maxc1,
|
||||
int qrow0, int qrow1,
|
||||
int mask_b_stride, int mask_q_stride,
|
||||
int mask_batch,
|
||||
int mask_b_stride, int mask_h_stride, int mask_q_stride,
|
||||
int mask_batch, int mask_head,
|
||||
const bool* __restrict__ mask,
|
||||
float Sacc[Traits::NC8][4],
|
||||
float Oacc[Traits::DN8][4],
|
||||
@@ -204,8 +204,8 @@ __device__ inline void mma_softmax_tile(
|
||||
int tid4 = lane & 3;
|
||||
|
||||
float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX;
|
||||
int mask_base0 = mask_batch * mask_b_stride + qrow0 * mask_q_stride;
|
||||
int mask_base1 = mask_batch * mask_b_stride + qrow1 * mask_q_stride;
|
||||
int mask_base0 = mask_batch * mask_b_stride + mask_head * mask_h_stride + qrow0 * mask_q_stride;
|
||||
int mask_base1 = mask_batch * mask_b_stride + mask_head * mask_h_stride + qrow1 * mask_q_stride;
|
||||
#pragma unroll
|
||||
for (int n8 = 0; n8 < Traits::NC8; n8++) {
|
||||
int cc = kv0 + n8 * 8 + 2 * tid4;
|
||||
|
||||
@@ -31,7 +31,7 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
|
||||
int ch_begin = split * chunks_per_split;
|
||||
int ch_end = min(chunks_total, ch_begin + chunks_per_split);
|
||||
|
||||
const int mask_base = batch * p.mask_b_stride;
|
||||
const int mask_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
|
||||
|
||||
for (int ci = ch_begin; ci < ch_end; ci++) {
|
||||
int chunk_start = ci * PDC_CHUNK;
|
||||
@@ -77,8 +77,8 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
|
||||
}
|
||||
|
||||
float new_m = fmaxf(m, partial);
|
||||
float alpha = expf(m - new_m);
|
||||
float beta = expf(partial - new_m);
|
||||
float alpha = __expf(m - new_m);
|
||||
float beta = __expf(partial - new_m);
|
||||
d = d * alpha + beta;
|
||||
|
||||
int pos = chunk_start + s;
|
||||
@@ -133,8 +133,8 @@ __global__ void paged_attn_decode_combine_kernel(PagedAttentionParams<bf16> p) {
|
||||
if (mi <= -FLT_MAX) continue;
|
||||
float li = mlp[s * 2 + 1];
|
||||
float nm = fmaxf(m, mi);
|
||||
float corr = expf(m - nm);
|
||||
float e = expf(mi - nm);
|
||||
float corr = __expf(m - nm);
|
||||
float e = __expf(mi - nm);
|
||||
acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
|
||||
l = fmaf(l, corr, li * e);
|
||||
m = nm;
|
||||
|
||||
@@ -55,19 +55,21 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
|
||||
const int64_t head_off = (int64_t)kv_head * Traits::HEAD_DIM;
|
||||
|
||||
// ---- Load tile lambda: paged addressing ----
|
||||
// Unified per-element page-table lookup. When page_size >= BC, all
|
||||
// elements in a tile share the same page, so the lookup is redundant
|
||||
// but harmless (L1-cached). This avoids a branch on page_size.
|
||||
auto load_tile = [&](int ti, int buf) {
|
||||
int kv0 = ti * Traits::BC;
|
||||
bf16* dK = sK + buf * Traits::BC * Traits::LD;
|
||||
bf16* dV = sV + buf * Traits::BC * Traits::LD;
|
||||
int logical_page = kv0 / p.page_size;
|
||||
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
||||
bool page_valid = (phys_page >= 0);
|
||||
#pragma unroll
|
||||
for (int i = lane * Traits::VEC; i < Traits::TOTAL;
|
||||
i += Traits::NUM_THREADS * Traits::VEC) {
|
||||
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
||||
int kc = kv0 + r;
|
||||
bool valid = (kc < p.kv_len) && page_valid;
|
||||
bool valid = (kc < p.kv_len);
|
||||
int phys_page = valid ? p.page_table[batch * p.max_pages + kc] : 0;
|
||||
valid = valid && (phys_page >= 0);
|
||||
int page_off = kc % p.page_size;
|
||||
int64_t gmem_base = (int64_t)phys_page * page_stride
|
||||
+ (int64_t)page_off * pos_stride
|
||||
@@ -110,8 +112,8 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
|
||||
int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
||||
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
|
||||
0, 0,
|
||||
p.mask_b_stride, 0,
|
||||
batch,
|
||||
p.mask_b_stride, 0, 0,
|
||||
batch, 0,
|
||||
p.mask,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
|
||||
|
||||
@@ -64,7 +64,7 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
||||
|
||||
// KV: stride-based base
|
||||
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||
int mask_batch_base = batch * p.mask_b_stride;
|
||||
int mask_batch_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
|
||||
int tiles = (p.kv_len + P_BC - 1) / P_BC;
|
||||
int tt = G * ROWS;
|
||||
int lid = row * G + gpos;
|
||||
|
||||
@@ -114,8 +114,8 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
|
||||
: p.kv_len;
|
||||
mma_softmax_tile<Traits, HasMask>(kv0, maxc0, maxc1,
|
||||
qr0, qr1,
|
||||
p.mask_b_stride, p.mask_q_stride,
|
||||
batch,
|
||||
p.mask_b_stride, p.mask_h_stride, p.mask_q_stride,
|
||||
batch, q_head,
|
||||
p.mask,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
|
||||
|
||||
@@ -103,6 +103,7 @@ inline void set_default_strides(P& p) {
|
||||
p.kv_stride_l = p.head_dim;
|
||||
p.kv_stride_d = 1;
|
||||
p.mask_b_stride = p.kv_len;
|
||||
p.mask_h_stride = 0;
|
||||
p.mask_q_stride = 0;
|
||||
}
|
||||
|
||||
@@ -114,6 +115,7 @@ inline void set_default_paged_strides(P& p) {
|
||||
p.q_stride_l = p.head_dim;
|
||||
p.q_stride_d = 1;
|
||||
p.mask_b_stride = p.kv_len;
|
||||
p.mask_h_stride = 0;
|
||||
p.mask_q_stride = 0;
|
||||
}
|
||||
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
<div align="center">
|
||||
|
||||
<img src="../images/logo.png" width="auto" alt="Logo">
|
||||
<img src="./images/logo.png" width="auto" alt="Logo">
|
||||
|
||||
<div>
|
||||
<a href="../../README.md">English</a> •
|
||||
<a href="../README.md">English</a> •
|
||||
<a href="#chinese">中文</a>
|
||||
</div>
|
||||
|
||||
@@ -23,7 +23,7 @@
|
||||
<br>
|
||||
|
||||
<div align="center">
|
||||
<a href="../../README.md">English</a> •
|
||||
<a href="../README.md">English</a> •
|
||||
<a href="#chinese">中文</a> •
|
||||
<a href="https://github.com/ViperEkura/AstrAI/issues">问题追踪</a> •
|
||||
<a href="https://github.com/ViperEkura/AstrAI/discussions">讨论区</a> •
|
||||
@@ -219,18 +219,23 @@ curl -X POST http://localhost:8000/v1/messages \
|
||||
curl http://localhost:8000/health
|
||||
```
|
||||
|
||||
SSE 流式格式、错误码和统计端点详见[推理文档](./inference.md)。
|
||||
SSE 流式格式、错误码和统计端点详见[推理文档](guides/inference.md)。
|
||||
|
||||
### 文档
|
||||
|
||||
| 文档 | 说明 |
|
||||
|------|------|
|
||||
| [CLI 参考](./params.md) | 所有 CLI 工具参数(训练、服务、生成、预处理) |
|
||||
| [架构文档](./architecture.md) | 系统架构、类图与设计模式 |
|
||||
| [训练文档](./training.md) | 训练循环、策略与公式 |
|
||||
| [推理文档](./inference.md) | KVCache、连续批处理、采样与 HTTP API |
|
||||
| [数据流程](./dataflow.md) | 数据管道、存储后端与数据集架构 |
|
||||
| [数据预处理](./preprocessing.md) | 声明式 JSON 驱动数据预处理 |
|
||||
| [快速上手](./get-started.md) | 安装与快速入门 |
|
||||
| [CLI 参考](./guides/params.md) | 所有 CLI 工具参数(训练、服务、生成、预处理) |
|
||||
| [数据预处理](./guides/preprocessing.md) | 声明式 JSON 驱动数据预处理 |
|
||||
| [训练文档](./guides/training.md) | 训练循环、策略与公式 |
|
||||
| [推理文档](./guides/inference.md) | KVCache、连续批处理、采样与 HTTP API |
|
||||
| [评估文档](./guides/evaluation.md) | HumanEval、MMLU、PPL、ROUGE、IFD、IFEval |
|
||||
| [分布式训练](./guides/distributed.md) | 多卡 DDP / FSDP 训练 |
|
||||
| [架构文档](./developer/architecture.md) | 系统架构、类图与设计模式 |
|
||||
| [数据流程](./developer/dataflow.md) | 数据管道、存储后端与数据集架构 |
|
||||
| [内部实现](./developer/internals.md) | 训练原理:损失公式、回调生命周期、KV Cache |
|
||||
| [CUDA 内核](./developer/cuda_kernels.md) | 自定义 CUDA 注意力内核与基准测试 |
|
||||
|
||||
### 贡献
|
||||
|
||||
@@ -277,7 +277,7 @@ classDiagram
|
||||
+ModuleList layers
|
||||
+RMSNorm norm
|
||||
+Linear lm_head
|
||||
+forward(input_ids, input_mask, paged_cache, position_ids) Dict[str, Tensor]
|
||||
+forward(input_ids, input_mask, kv_cache, position_ids) Dict[str, Tensor]
|
||||
+load_state_dict(state_dict, strict, assign)
|
||||
+state_dict()
|
||||
}
|
||||
@@ -299,7 +299,7 @@ classDiagram
|
||||
+RMSNorm input_norm
|
||||
+nn.Module mlp # MLP or DeepSeekMoE via FFNFactory
|
||||
+RMSNorm post_attention_norm
|
||||
+forward(x, rotary_emb, attention_mask, paged_cache) Tensor
|
||||
+forward(x, rotary_emb, attention_mask, kv_cache) Tensor
|
||||
}
|
||||
|
||||
class GQA {
|
||||
@@ -314,7 +314,7 @@ classDiagram
|
||||
+Linear q_proj, k_proj, v_proj, o_proj
|
||||
+Linear gate # only if use_gated_attention
|
||||
+RMSNorm q_norm, k_norm # only if use_qk_norm
|
||||
+forward(x, rotary_emb, attn_mask, paged_cache) Tensor
|
||||
+forward(x, rotary_emb, attn_mask, kv_cache) Tensor
|
||||
}
|
||||
|
||||
class MLA {
|
||||
@@ -334,7 +334,7 @@ classDiagram
|
||||
+Linear gate # only if use_gated_attention
|
||||
+RMSNorm kv_norm
|
||||
+RMSNorm q_norm, k_norm # only if use_qk_norm
|
||||
+forward(x, rotary_emb, attn_mask, paged_cache) Tensor
|
||||
+forward(x, rotary_emb, attn_mask, kv_cache) Tensor
|
||||
}
|
||||
|
||||
class MLP {
|
||||
@@ -824,75 +824,46 @@ classDiagram
|
||||
+record(page_idx, token_ids, logical_page_idx)
|
||||
}
|
||||
|
||||
class Storage {
|
||||
+int page_size
|
||||
+Tensor k_cache
|
||||
+Tensor v_cache
|
||||
+write(layer_id, page_table, start_pos, k, v)
|
||||
+gather(layer_id, page_table, total_len) Tuple[Tensor, Tensor]
|
||||
class KVStorage {
|
||||
+int size
|
||||
+Tensor k_buffer
|
||||
+Tensor v_buffer
|
||||
+get_key_buffer(layer_id) Tensor
|
||||
+get_value_buffer(layer_id) Tensor
|
||||
+set_kv_buffer(layer_id, loc, k, v)
|
||||
}
|
||||
|
||||
class ReqToTokenPool {
|
||||
+int size
|
||||
+int max_context_len
|
||||
+Tensor req_to_token
|
||||
+alloc(num_reqs) List[int]
|
||||
+free(req_indices)
|
||||
+write(indices, values)
|
||||
}
|
||||
|
||||
class KVCache {
|
||||
<<abstract>>
|
||||
+task_alloc(task_id, prompt_ids) bool
|
||||
+task_free(task_id)
|
||||
+task_extend(task_id, pos) bool
|
||||
+task_cached(task_id) int
|
||||
+task_record_hashes(task_id, prompt_ids, start_logical_page)
|
||||
+bind_tasks(task_ids, total_len, device) CacheView
|
||||
+Tensor k_buffer
|
||||
+Tensor v_buffer
|
||||
+Tensor req_to_token
|
||||
+Tensor req_pool_indices
|
||||
+Tensor seq_lens
|
||||
+Tensor out_cache_loc
|
||||
}
|
||||
|
||||
class PageCache {
|
||||
class PagePool {
|
||||
+int page_size
|
||||
-PagePool _pool
|
||||
-Storage _storage
|
||||
-TaskTable _table
|
||||
+bool contiguous
|
||||
-KVStorage _storage
|
||||
-ReqToTokenPool _req_pool
|
||||
-Allocator _alloc
|
||||
-PrefixCache _prefix
|
||||
+task_alloc(task_id, prompt_ids) bool
|
||||
+task_free(task_id)
|
||||
+task_extend(task_id, pos) bool
|
||||
+task_cached(task_id) int
|
||||
+task_record_hashes(task_id, prompt_ids, start_logical_page)
|
||||
+bind_tasks(task_ids, total_len, device) PageCacheView
|
||||
}
|
||||
|
||||
class ContiguousCache {
|
||||
+int max_seq_len
|
||||
+Tensor k, v
|
||||
+task_alloc(task_id, prompt_ids) bool
|
||||
+task_free(task_id)
|
||||
+task_extend(task_id, pos) bool
|
||||
+bind_tasks(task_ids, total_len, device) ContiguousCacheView
|
||||
}
|
||||
|
||||
class CacheView {
|
||||
<<abstract>>
|
||||
+write(layer_id, k, v)
|
||||
+gather(layer_id) Tuple[Tensor, Tensor]
|
||||
}
|
||||
|
||||
class PageCacheView {
|
||||
-Storage _storage
|
||||
+Tensor _page_table
|
||||
+int _total_len
|
||||
+write(layer_id, k, v)
|
||||
+gather(layer_id) Tuple[Tensor, Tensor]
|
||||
}
|
||||
|
||||
class ContiguousCacheView {
|
||||
-ContiguousCache _cache
|
||||
+Tensor _batch_indices
|
||||
+int _total_len
|
||||
+write(layer_id, k, v)
|
||||
+gather(layer_id) Tuple[Tensor, Tensor]
|
||||
}
|
||||
|
||||
class TaskTable {
|
||||
+set(task_id, page_table, cached)
|
||||
+get(task_id) List[int]
|
||||
+get_cached(task_id) int
|
||||
+get_ref(task_id) List[int]
|
||||
+pop(task_id) Tuple[List[int], int]
|
||||
+table_tensor(task_ids, device) Tensor
|
||||
+bind_tasks(task_ids, seq_lens, device, start_pos) KVCache
|
||||
}
|
||||
|
||||
class Task {
|
||||
@@ -924,7 +895,6 @@ classDiagram
|
||||
+AutoTokenizer tokenizer
|
||||
+int max_batch_size
|
||||
+int max_seq_len
|
||||
+int max_prompt_len
|
||||
+Deque waiting_queue
|
||||
+List active_tasks
|
||||
+add_task(prompt, max_tokens, temperature, top_p, top_k, stream_callback) str
|
||||
@@ -1196,11 +1166,6 @@ classDiagram
|
||||
}
|
||||
|
||||
class FSDPExecutor {
|
||||
-_prepare_model(model) nn.Module
|
||||
+unwrap_model(model) dict
|
||||
}
|
||||
|
||||
class FSDP2Executor {
|
||||
-_prepare_model(model) nn.Module
|
||||
-_no_sync(model) context manager
|
||||
+unwrap_model(model) dict
|
||||
@@ -1303,7 +1268,6 @@ classDiagram
|
||||
BaseExecutor <|-- NoneExecutor
|
||||
BaseExecutor <|-- DDPExecutor
|
||||
BaseExecutor <|-- FSDPExecutor
|
||||
BaseExecutor <|-- FSDP2Executor
|
||||
ResponseBuilder <|-- OpenAIResponseBuilder
|
||||
ResponseBuilder <|-- AnthropicResponseBuilder
|
||||
BaseToolParser <|-- SimpleJsonToolParser
|
||||
@@ -1321,17 +1285,13 @@ classDiagram
|
||||
RawRollout <|-- RolloutResult
|
||||
LaunchStrategy <|-- TorchrunStrategy
|
||||
LaunchStrategy <|-- LocalStrategy
|
||||
KVCache <|-- PageCache
|
||||
KVCache <|-- ContiguousCache
|
||||
CacheView <|-- PageCacheView
|
||||
CacheView <|-- ContiguousCacheView
|
||||
|
||||
%% --- Composition (strong ownership, part destroyed with whole) ---
|
||||
PageCache *-- PagePool
|
||||
PageCache *-- Storage
|
||||
PageCache *-- TaskTable
|
||||
PagePool *-- KVStorage
|
||||
PagePool *-- ReqToTokenPool
|
||||
PagePool *-- Allocator
|
||||
PagePool *-- PrefixCache
|
||||
InferenceEngine *-- InferenceScheduler
|
||||
InferenceScheduler *-- KVCache
|
||||
InferenceScheduler *-- PagePool
|
||||
InferenceScheduler *-- Executor
|
||||
InferenceScheduler *-- TaskManager
|
||||
AutoRegressiveLM *-- DecoderBlock
|
||||
@@ -1359,8 +1319,6 @@ classDiagram
|
||||
TrainContext o-- BaseScheduler
|
||||
TrainContext o-- Checkpoint
|
||||
TrainContext o-- BaseExecutor
|
||||
PageCacheView o-- Storage
|
||||
ContiguousCacheView o-- ContiguousCache
|
||||
SamplingPipeline o-- BaseSamplingStrategy
|
||||
BaseDataset o-- Store
|
||||
Pipeline o-- PipelineConfig
|
||||
@@ -1397,7 +1355,6 @@ classDiagram
|
||||
ExecutorFactory ..> NoneExecutor : creates
|
||||
ExecutorFactory ..> DDPExecutor : creates
|
||||
ExecutorFactory ..> FSDPExecutor : creates
|
||||
ExecutorFactory ..> FSDP2Executor : creates
|
||||
ToolParserFactory ..> BaseToolParser : creates
|
||||
TrainContextBuilder ..> ExecutorFactory : creates
|
||||
Trainer ..> TrainContextBuilder : uses
|
||||
@@ -1406,8 +1363,7 @@ classDiagram
|
||||
TrainContextBuilder ..> RDSampler : creates
|
||||
Checkpoint ..> Checkpoint : serializes
|
||||
CheckpointCallback ..> Checkpoint : creates
|
||||
PageCache ..> PageCacheView : binds
|
||||
ContiguousCache ..> ContiguousCacheView : binds
|
||||
PagePool ..> KVCache : binds
|
||||
InferenceEngine ..> GenerationRequest : uses
|
||||
InferenceEngine ..> GenerateResult : creates
|
||||
OpenAIResponseBuilder ..> ChatCompletionRequest : receives
|
||||
@@ -1444,8 +1400,9 @@ classDiagram
|
||||
| **astrai.model** | AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, LoRAConfig, LoRALinear, RotaryEmbedding, Embedding | Neural network model |
|
||||
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
|
||||
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy–GRPOStrategy, StrategyFactory, BaseScheduler–WSDScheduler, SchedulerFactory, TrainCallback(Protocol)–MetricCallback, CallbackFactory, RawRollout, RolloutResult, BaseRewardModel, RolloutGenerator, RolloutRunner | Training workflow |
|
||||
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, KVCache–ContiguousCache/PageCache, CacheView–ContiguousCacheView/PageCacheView, Allocator–Storage, Task, TaskManager, TaskStatus, StreamDecoder, GenerationRequest, GenerateResult, BaseSamplingStrategy–SamplingPipeline, FrequencyPenaltyStrategy, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, StopInfo, ChatMessage, FunctionDef, ToolDef, ChatCompletionRequest, AnthropicMessage, MessagesRequest, BaseToolParser, ToolParserFactory, SimpleJsonToolParser | Inference service |
|
||||
| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, LaunchStrategy, TorchrunStrategy, LocalStrategy, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, FSDPExecutor, FSDP2Executor, GradientState, AccumOptimizer, AccumScheduler, ParallelModel, RowParallelLinear, ColumnParallelLinear | Distributed parallel & gradient accumulation |
|
||||
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, PagePool, KVStorage, ReqToTokenPool, KVCache, Allocator, PrefixCache, Task, TaskManager, TaskStatus, StreamDecoder, GenerationRequest, GenerateResult, BaseSamplingStrategy–SamplingPipeline, FrequencyPenaltyStrategy, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, StopInfo, ChatMessage, FunctionDef, ToolDef, ChatCompletionRequest, AnthropicMessage, MessagesRequest, BaseToolParser, ToolParserFactory, SimpleJsonToolParser | Inference service |
|
||||
| **astrai.extension** | AttentionBackend, TorchNativeBackend, CudaBackend, attn_backend, ATTN_BACKEND, attn_decode, attn_prefill, attn_paged_decode, is_available | CUDA attention kernels + backend abstraction |
|
||||
| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, LaunchStrategy, TorchrunStrategy, LocalStrategy, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, FSDPExecutor, GradientState, AccumOptimizer, AccumScheduler | Distributed parallel & gradient accumulation |
|
||||
| **astrai.factory** | BaseFactory | Component registration |
|
||||
| **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers |
|
||||
|
||||
@@ -1462,7 +1419,8 @@ classDiagram
|
||||
| **Observer** | `TrainCallback`, callback implementations | Training process monitoring |
|
||||
| **Context** | `TrainContext` | Unified training state bag |
|
||||
| **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction |
|
||||
| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor`, `FSDP2Executor` | Gradient accumulation & model distribution |
|
||||
| **Strategy (Attention)** | `AttentionBackend`, `TorchNativeBackend`, `CudaBackend` | Attention computation backend switching via context manager |
|
||||
| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor` | Gradient accumulation & model distribution |
|
||||
| **Storage** | `Store`, `H5Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support |
|
||||
| **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching |
|
||||
| **AutoModel Registry** | `AutoModel`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading |
|
||||
@@ -1472,8 +1430,8 @@ classDiagram
|
||||
1. **Config → Training**: `TrainConfig` holds `model_fn`, `dataset`, `optimizer_fn`, `scheduler_fn`, `parallel_mode`, `executor_kwargs`
|
||||
2. **Training Flow**: `Trainer` → `TrainContextBuilder` → `TrainContext`, uses `BaseStrategy` for loss, `BaseExecutor` for gradient accumulation + model distribution
|
||||
3. **Strategy Selection**: `StrategyFactory` creates strategy by `train_type`
|
||||
4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)` → `NoneExecutor` / `DDPExecutor` / `FSDPExecutor` / `FSDP2Executor`
|
||||
5. **Inference Flow**: `InferenceEngine` → `InferenceScheduler` → `AutoRegressiveLM`, backed by `KVCache` + `SamplingPipeline`
|
||||
4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)` → `NoneExecutor` / `DDPExecutor` / `FSDPExecutor`
|
||||
5. **Inference Flow**: `InferenceEngine` → `InferenceScheduler` → `AutoRegressiveLM`, backed by `PagePool` + `KVCache` + `SamplingPipeline`. Attention backend selected via `attn_backend()` context manager (`TorchNativeBackend` default, `CudaBackend` for CUDA kernels).
|
||||
6. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP
|
||||
7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (H5Store/MmapStore/JsonlStore) loads data with explicit `_length` and multi-segment `_data`
|
||||
8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata (rank-0 only), extra state saved as `{key}.pt`
|
||||
@@ -1481,4 +1439,4 @@ classDiagram
|
||||
10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops
|
||||
11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers
|
||||
|
||||
> Document Update Time: 2026-07-20
|
||||
> Document Update Time: 2026-07-30
|
||||
@@ -0,0 +1,148 @@
|
||||
# CUDA Kernels
|
||||
|
||||
AstrAI includes optional custom CUDA attention kernels for decode and prefill. These are built when `nvcc` is available and CUDA is detected, and are dispatched via the `CudaBackend` attention backend.
|
||||
|
||||
## Overview
|
||||
|
||||
| Kernel | File | Description |
|
||||
|--------|------|-------------|
|
||||
| `attn_decode` | `attn_decode.cu` | GQA decode attention (split-KV) |
|
||||
| `attn_prefill` | `attn_prefill.cu` | GQA prefill attention (split-Q) |
|
||||
| `attn_paged_decode` | `attn_paged_decode.cu` | Paged KV cache decode attention |
|
||||
|
||||
Additionally, optimized `.cuh` variants with tensor-core MMA (Matrix Multiply-Accumulate) exist:
|
||||
|
||||
| Variant | File | Optimization |
|
||||
|---------|------|--------------|
|
||||
| Split-KV MMA decode | `attn_decode_split_kv_mma.cuh` | Split KV across warps + MMA (sm_80+) |
|
||||
| Split-Q MMA prefill | `attn_prefill_split_q_mma.cuh` | Split Q across warps + MMA (sm_80+) |
|
||||
| Paged split-KV MMA decode | `attn_paged_decode_split_kv_mma.cuh` | Paged cache + split-KV + MMA |
|
||||
|
||||
## Build System
|
||||
|
||||
### Auto-detection
|
||||
|
||||
Kernels are built when **both** of these conditions are met:
|
||||
1. `nvcc` is available on `PATH`
|
||||
2. `torch.cuda.is_available()` returns `True`
|
||||
|
||||
Unless `CSRC_KERNELS=false` is set explicitly.
|
||||
|
||||
### Manual build
|
||||
|
||||
```bash
|
||||
# During install
|
||||
CSRC_KERNELS=true pip install -e . --no-build-isolation
|
||||
|
||||
# Rebuild after editing .cu/.cuh files
|
||||
CSRC_KERNELS=true python setup.py build_ext --inplace
|
||||
# Output: astrai/extension/*.so
|
||||
```
|
||||
|
||||
### Architecture flags
|
||||
|
||||
`csrc/build.py` auto-detects the GPU compute capability and generates the appropriate `nvcc` gencode flag:
|
||||
|
||||
- **sm_80+** (Ampere and later): enables tensor-core MMA path (`mma.sync.m16n8k16.bf16`)
|
||||
- **Below sm_80**: adds `-DASTRAI_NO_MMA` to disable the MMA path at compile time
|
||||
|
||||
### Build configuration
|
||||
|
||||
```
|
||||
NVCC_FLAGS = -O3 --expt-relaxed-constexpr --use_fast_math
|
||||
--ptxas-options=-O3,-v --extra-device-vectorization --threads=8
|
||||
```
|
||||
|
||||
The `REGISTRY` in `csrc/build.py` lists all registered kernels (currently 3). Each entry maps a kernel name to its source files and build flags.
|
||||
|
||||
## Attention Backend
|
||||
|
||||
`astrai/extension/attention_backend.py` provides the backend abstraction:
|
||||
|
||||
- **`AttentionBackend`** (ABC): `fwd_decode` / `fwd_prefill` abstract methods, `forward` dispatches by q_len
|
||||
- **`TorchNativeBackend`**: SDPA with indirect KV cache gather (default)
|
||||
- **`CudaBackend`**: CUDA kernel dispatch — decode via `attn_paged_decode` (page_size=1), prefill via `attn_prefill`
|
||||
|
||||
Select a backend via context manager (mirrors `torch.nn.attention.sdpa_kernel`):
|
||||
|
||||
```python
|
||||
from astrai.extension import attn_backend, ATTN_BACKEND
|
||||
|
||||
with attn_backend(ATTN_BACKEND.CUDA):
|
||||
engine.generate("hello")
|
||||
```
|
||||
|
||||
`CudaBackend` falls back to `TorchNativeBackend` when a kernel is not available.
|
||||
|
||||
## Python Wrappers
|
||||
|
||||
`astrai/extension/attention_ops.py` provides Python wrappers for each compiled kernel. Each wrapper calls its CUDA kernel directly and raises `RuntimeError` if the `.so` is not available. Fallback to torch SDPA is handled by the attention backend, not the wrapper functions.
|
||||
|
||||
Interface (all functions):
|
||||
```
|
||||
is_causal: True = causal mask; False = non-causal
|
||||
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
|
||||
```
|
||||
|
||||
Layout convention: all q/k/v are `[batch, seq_len, n_heads, head_dim]` (blhd). Scale is always `1/sqrt(head_dim)`.
|
||||
|
||||
## Standalone Testing
|
||||
|
||||
Each `csrc/tests/*.cu` file has the `nvcc` compile command in its header comment. Example:
|
||||
|
||||
```bash
|
||||
nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \
|
||||
--ptxas-options=-O3,-v --extra-device-vectorization \
|
||||
csrc/tests/attn_decode_test.cu -o /tmp/test && /tmp/test
|
||||
```
|
||||
|
||||
Test files:
|
||||
- `attn_decode_test.cu` — basic decode kernel
|
||||
- `attn_paged_decode_test.cu` — paged decode kernel
|
||||
- `attn_prefill_test.cu` — prefill kernel
|
||||
|
||||
## Benchmarks
|
||||
|
||||
Hardware: NVIDIA L20 (sm_89, 46 GB), CUDA 12.8, driver 570.86.
|
||||
|
||||
Reproduce:
|
||||
```bash
|
||||
nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \
|
||||
--ptxas-options=-O3,-v --extra-device-vectorization \
|
||||
csrc/tests/attn_<name>_test.cu -o /tmp/test && /tmp/test
|
||||
```
|
||||
|
||||
## Known Optimization Targets
|
||||
|
||||
- **Decode D=256**: spill eliminated (BC=16 + STAGES=2), but still 248 regs — further tiling could help.
|
||||
- **Prefill single-batch**: bandwidth low (52 GB/s at q=kv=2048) — likely compute-bound but near L20 bf16 ceiling (~94 TFLOP/s).
|
||||
- **Decode single-batch**: bandwidth low (309 GB/s at kv=512) — L20 HBM ~864 GB/s theoretical; small kv underutilizes SMs despite split-KV.
|
||||
|
||||
## File Layout
|
||||
|
||||
```
|
||||
csrc/
|
||||
├── build.py # Build system: REGISTRY, _arch_flags, nvcc flags
|
||||
├── kernels/
|
||||
│ ├── attn_common.h # Shared attention utilities
|
||||
│ ├── attn_decode.cu # Basic decode kernel (registered)
|
||||
│ ├── attn_prefill.cu # Basic prefill kernel (registered)
|
||||
│ ├── attn_paged_decode.cu # Paged decode kernel (registered)
|
||||
│ ├── attn_decode_split_kv.cuh # Split-KV variant
|
||||
│ ├── attn_decode_split_kv_mma.cuh # Split-KV + MMA variant
|
||||
│ ├── attn_prefill_split_q.cuh # Split-Q variant
|
||||
│ ├── attn_prefill_split_q_mma.cuh # Split-Q + MMA variant
|
||||
│ ├── attn_paged_decode_split_kv.cuh # Paged + split-KV variant
|
||||
│ ├── attn_paged_decode_split_kv_mma.cuh # Paged + split-KV + MMA variant
|
||||
│ ├── attn_dispatchers.cuh # Kernel dispatch macros
|
||||
│ ├── attn_entry_utils.cuh # Entry point helpers
|
||||
│ ├── attn_mma_utils.cuh # MMA utilities
|
||||
│ └── attn_warp_utils.cuh # Warp-level utilities
|
||||
└── tests/
|
||||
├── test_utils.cuh # Shared test utilities
|
||||
├── attn_decode_test.cu # Decode kernel test
|
||||
├── attn_paged_decode_test.cu # Paged decode test
|
||||
└── attn_prefill_test.cu # Prefill kernel test
|
||||
```
|
||||
|
||||
> Document Update Time: 2026-07-30
|
||||
@@ -1,6 +1,6 @@
|
||||
# Data Flow
|
||||
|
||||
This document describes the data pipeline: from raw text to model input tensors. For creating preprocessing configs, see [Preprocessing Guide](preprocessing.md).
|
||||
This document describes the data pipeline: from raw text to model input tensors. For creating preprocessing configs, see [Preprocessing Guide](../guides/preprocessing.md).
|
||||
|
||||
## Contents
|
||||
|
||||
@@ -33,7 +33,7 @@ Raw text is tokenized via `AutoTokenizer.encode()` and saved as HDF5 (`.h5`) or
|
||||
|
||||
### Tokenization
|
||||
|
||||
The `Pipeline` reads JSONL lines, applies the mask builder (see [Preprocessing](preprocessing.md)), and produces flat token sequences:
|
||||
The `Pipeline` reads JSONL lines, applies the mask builder (see [Preprocessing](../guides/preprocessing.md)), and produces flat token sequences:
|
||||
|
||||
```python
|
||||
# Per JSONL line: messages → chat template → token IDs + loss mask
|
||||
@@ -83,9 +83,11 @@ All backends normalise tensors into `Store._data[Dict[str, List[Tensor]]]` + `St
|
||||
## Dataset Architecture
|
||||
|
||||
```
|
||||
DatasetFactory.load(train_type, load_path, window_size, stride=None,
|
||||
storage_type=None, tokenizer_path=None,
|
||||
max_position_embeddings=2048, store=None)
|
||||
DatasetFactory.load(
|
||||
train_type, load_path=None, window_size=0, stride=None,
|
||||
storage_type=None, tokenizer_path=None,
|
||||
max_len=2048, store=None
|
||||
)
|
||||
→ BaseDataset.load(load_path, storage_type=None)
|
||||
→ detect_format(load_path)
|
||||
→ StoreFactory.create(storage_type)
|
||||
@@ -0,0 +1,231 @@
|
||||
# Internals
|
||||
|
||||
Mathematical foundations and internal algorithms for AstrAI's training, inference, and preprocessing pipelines. For practical usage guides, see [Training](../guides/training.md), [Inference](../guides/inference.md), and [Preprocessing](../guides/preprocessing.md).
|
||||
|
||||
## Contents
|
||||
|
||||
- [Autoregression & Causal Masking](#autoregression--causal-masking)
|
||||
- [Rotary Position Embedding (RoPE)](#rotary-position-embedding-rope)
|
||||
- [Training Loss Formulas](#training-loss-formulas)
|
||||
- [Training Loop Internals](#training-loop-internals)
|
||||
- [Callback Lifecycle](#callback-lifecycle)
|
||||
- [KV Cache Mathematics](#kv-cache-mathematics)
|
||||
- [Mask Algorithm Internals](#mask-algorithm-internals)
|
||||
- [Gradient Accumulation Mechanics](#gradient-accumulation-mechanics)
|
||||
|
||||
## Autoregression & Causal Masking
|
||||
|
||||
Given a token sequence, the model predicts the probability of the next token. Each generated token is appended to the input and fed back, repeating until an end-of-sequence token or max length.
|
||||
|
||||
```
|
||||
sequence : [[1, 2, 3, 4, 5, 6]]
|
||||
input_ids: [[1, 2, 3, 4, 5]]
|
||||
target_ids: [[2, 3, 4, 5, 6]]
|
||||
```
|
||||
|
||||
A lower-triangular causal mask prevents attending to future positions:
|
||||
|
||||
```
|
||||
[[0, -inf, -inf, -inf, -inf],
|
||||
[0, 0, -inf, -inf, -inf],
|
||||
[0, 0, 0, -inf, -inf],
|
||||
[0, 0, 0, 0, -inf],
|
||||
[0, 0, 0, 0, 0]]
|
||||
```
|
||||
|
||||
This ensures position $i$ can only attend to positions $\leq i$, which is essential for autoregressive generation.
|
||||
|
||||
## Rotary Position Embedding (RoPE)
|
||||
|
||||
RoPE embeds position into Q/K vectors via complex rotation:
|
||||
|
||||
$$ q_i = R_i W_q x_i, \quad k_j = R_j W_k x_j, \quad q_i^T k_j = x_i^T W_q^T R_{i-j} W_k x_j $$
|
||||
|
||||
The complex rotation `freqs_cis` is pre-computed once (`cos, sin` pairs per position). `apply_rotary_emb` multiplies Q/K as complex numbers. The key property is that the dot product $q_i^T k_j$ depends only on the relative position $i - j$, not the absolute positions.
|
||||
|
||||
**Critical for inference**: RoPE is applied **before** KV cache write, not after. If applied after caching, position encoding drift occurs because cached K/V would have stale rotation factors.
|
||||
|
||||
## Training Loss Formulas
|
||||
|
||||
### SEQ (Pre-training)
|
||||
|
||||
Next-token cross-entropy with optional label smoothing:
|
||||
|
||||
$$ L_{\text{PT}} = -\sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta) $$
|
||||
|
||||
### SFT (Supervised Fine-Tuning)
|
||||
|
||||
Masked cross-entropy (`ignore_index=-100`) over response tokens only:
|
||||
|
||||
$$ L_{\text{SFT}} = -\sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta) $$
|
||||
|
||||
Prompt tokens are masked out via `loss_mask`; only response tokens contribute to the loss.
|
||||
|
||||
### DPO (Direct Preference Optimization)
|
||||
|
||||
Frozen reference model, preference margin via log-ratio:
|
||||
|
||||
$$ L_{\text{DPO}} = -\mathbb{E}\left[\log\sigma\left(\beta\log\frac{\pi_\theta(y_w\mid x)}{\pi_{\text{ref}}(y_w\mid x)} - \beta\log\frac{\pi_\theta(y_l\mid x)}{\pi_{\text{ref}}(y_l\mid x)}\right)\right] $$
|
||||
|
||||
Parameters: `beta=0.1`, `reduction="sum"`.
|
||||
|
||||
### GRPO (Group Relative Policy Optimization)
|
||||
|
||||
Token-level PPO with group-normalized advantages:
|
||||
|
||||
$$ \text{Advantage}_i = \frac{r_i - \mu}{\sigma + \epsilon} $$
|
||||
|
||||
$$ L_{\text{GRPO}} = -\mathbb{E}_t\left[\min\left(\rho_t A,\; \text{clip}\left(\rho_t, 1-\epsilon, 1+\epsilon\right)A\right)\right] + \lambda \cdot \mathbb{E}_t\left[\frac{\pi_{\text{ref}}}{\pi_\theta} - \log\frac{\pi_{\text{ref}}}{\pi_\theta} - 1\right] $$
|
||||
|
||||
Where $\rho_t = \pi_\theta(a_t|s_t) / \pi_{\text{old}}(a_t|s_t)$ is the per-token importance sampling ratio. Advantages are derived from scalar per-response rewards, group-normalized, and broadcast across all response tokens. Only response tokens contribute to the loss.
|
||||
|
||||
Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`.
|
||||
|
||||
## Training Loop Internals
|
||||
|
||||
Two-level loop: **epoch** → **batch**. Optimizer step fires every `grad_accum_steps` batches.
|
||||
|
||||
```
|
||||
on_train_begin
|
||||
model.train()
|
||||
on_epoch_begin
|
||||
for batch in dataloader:
|
||||
on_batch_begin
|
||||
with executor.accumulate(model):
|
||||
loss = strategy.compute_loss(batch)
|
||||
context.loss = loss.item()
|
||||
stand_loss = loss / executor.grad_accum_steps
|
||||
executor.backward(stand_loss)
|
||||
context.consumed_samples += (
|
||||
context.config.batch_per_device * context.world_size
|
||||
)
|
||||
on_batch_end
|
||||
|
||||
if executor.sync_gradients:
|
||||
on_optimizer_step
|
||||
optimizer.step()
|
||||
optimizer.zero_grad()
|
||||
if scheduler:
|
||||
scheduler.step()
|
||||
on_epoch_end
|
||||
on_train_end
|
||||
```
|
||||
|
||||
The loss is divided by `grad_accum_steps` before `backward()`, so accumulated gradients sum to the correct mean.
|
||||
|
||||
## Callback Lifecycle
|
||||
|
||||
| Hook | Fires | Default callback |
|
||||
|------|-------|-----------------|
|
||||
| `on_train_begin` | Before training starts | `GradientCheckpointingCallback` |
|
||||
| `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` |
|
||||
| `on_batch_begin` | Every batch | — |
|
||||
| `on_optimizer_step` | Every accumulation window | `GradientClippingCallback`, `MetricCallback`, `ProgressBarCallback` |
|
||||
| `on_batch_end` | Every batch | `CheckpointCallback` |
|
||||
| `on_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` |
|
||||
| `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` |
|
||||
| `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricCallback`, `GradientCheckpointingCallback` |
|
||||
|
||||
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm), `gradient_clipping` (always registered; computes grad norm, clips only when `max_grad_norm` is not `None`).
|
||||
|
||||
## KV Cache Mathematics
|
||||
|
||||
At decode time, only the last query token matters. All previous K/V are cached to avoid recomputation:
|
||||
|
||||
$$ o_n = \sum_j \text{softmax}\left(\frac{q_n k_j}{\sqrt{d_k}}\right) v_j $$
|
||||
|
||||
The cache stores $k_j$ and $v_j$ for all previous positions. At each decode step, only $q_n$ (the current query) is computed fresh, and attention is computed against the cached K/V.
|
||||
|
||||
**RoPE ordering**: RoPE is applied to Q/K **before** writing to the KV cache. This is essential because:
|
||||
1. The cached K values already contain the rotation for their original positions.
|
||||
2. The new Q is rotated for its current position.
|
||||
3. The dot product $q_n^T k_j$ then correctly depends on $n - j$ (relative position).
|
||||
|
||||
If RoPE were applied after caching, the rotation factors would be inconsistent between cached and new tokens.
|
||||
|
||||
### Cache Architecture
|
||||
|
||||
Three-layer separation (SGLang-inspired):
|
||||
|
||||
- **KVStorage**: Flat token-level buffers `[n_layers, size, n_kv_heads, head_dim]`.
|
||||
- **ReqToTokenPool**: Index table `[req_idx, pos] → physical token slot`, shared across all layers.
|
||||
- **Allocator + PrefixCache**: Paged-mode slot allocation with ref-counting, LRU eviction, and hash-based prefix sharing.
|
||||
|
||||
`PagePool` orchestrates all three. In contiguous mode (default), `req_to_token` is a trivial linear mapping. In paged mode, slots are allocated on demand with prefix caching support. Attention layers access buffers directly via `KVCache` dataclass — no methods, no abstraction.
|
||||
|
||||
### Attention Backend
|
||||
|
||||
Attention computation is decoupled from the model via `AttentionBackend` ABC (`astrai/extension/attention_backend.py`):
|
||||
|
||||
- **`TorchNativeBackend`** (default): writes K/V to cache, gathers via `req_to_token` indirect indexing, calls `F.scaled_dot_product_attention`.
|
||||
- **`CudaBackend`**: decode path uses `attn_paged_decode` with `page_size=1` (the `req_to_token` table serves as the page table, each token slot is a single-token "page"); prefill path gathers K/V then calls `attn_prefill`. Falls back to `TorchNativeBackend` when kernel unavailable.
|
||||
|
||||
Backend selection is thread-safe via `contextvars`, mirroring `torch.nn.attention.sdpa_kernel`:
|
||||
|
||||
```python
|
||||
from astrai.extension import attn_backend, ATTN_BACKEND
|
||||
|
||||
with attn_backend(ATTN_BACKEND.CUDA):
|
||||
engine.generate("hello")
|
||||
```
|
||||
|
||||
Layout convention: all q/k/v are `[batch, seq_len, n_heads, head_dim]` (blhd). Scale is always `1/sqrt(head_dim)`.
|
||||
|
||||
## Mask Algorithm Internals
|
||||
|
||||
### Template mode (`template: true`)
|
||||
|
||||
1. Prepend BOS token (masked)
|
||||
2. For each message in the field's array:
|
||||
1. Render through `chat_template` for that single message
|
||||
2. Encode rendered text
|
||||
3. Apply mask rule for the message's role
|
||||
|
||||
### Non-template mode
|
||||
|
||||
Encode the field value as text. Mask value is 1 (train) or 0 (mask) per the section's `action`.
|
||||
|
||||
### Text config detection
|
||||
|
||||
When no section uses `template` and all sections have `action: "train"`, the builder omits `loss_mask` from the output — all tokens are trained.
|
||||
|
||||
### Position ID strategies
|
||||
|
||||
| Mode | Behavior |
|
||||
|------|----------|
|
||||
| `none` | No position IDs generated |
|
||||
| `doc_reset` | Reset position to 0 at each document boundary in packed sequences |
|
||||
| `continuous` | Continuous position IDs across packed documents |
|
||||
|
||||
Default is `doc_reset`, which ensures each document in a packed bin starts from position 0, preventing position encoding drift between unrelated documents.
|
||||
|
||||
## Gradient Accumulation Mechanics
|
||||
|
||||
Three cooperating layers enable gradient accumulation:
|
||||
|
||||
1. **`GradientState`** — tracks the micro-step counter. Fires `sync_gradients=True` every `grad_accum_steps` micro-batches. The counter is incremented at the **start** of `accumulate()`, before the forward pass.
|
||||
|
||||
2. **`executor._no_sync(model)`** — suppresses gradient synchronization on non-sync micro-steps:
|
||||
- `NoneExecutor`: `nullcontext` (nothing to skip)
|
||||
- `DDPExecutor`: `model.no_sync()` (PyTorch's built-in — skips all-reduce of gradient buckets)
|
||||
- `FSDPExecutor`: `set_requires_gradient_sync(False, recurse=True)` on each `FSDPModule` (FSDP2's native mechanism)
|
||||
|
||||
3. **`AccumOptimizer` / `AccumScheduler`** — wrap the real optimizer/scheduler. `step()` and `zero_grad()` are gated on `sync_gradients` — they only forward to the inner optimizer when the sync flag is True.
|
||||
|
||||
The loss is divided by `grad_accum_steps` before `backward()`, so gradients sum to the correct mean across micro-steps. `consumed_samples` increments by `batch_per_device * world_size` every micro-batch.
|
||||
|
||||
### Effective batch size
|
||||
|
||||
$$ \text{Effective batch} = \text{nprocs} \times \text{batch\_per\_device} \times \text{grad\_accum\_steps} $$
|
||||
|
||||
### Total optimizer steps
|
||||
|
||||
```
|
||||
samples_per_replica = ceil(dataset_len / nprocs)
|
||||
batches_per_replica = ceil(samples_per_replica / batch_per_device)
|
||||
total_steps = (batches_per_replica // grad_accum_steps) * n_epoch
|
||||
```
|
||||
|
||||
This accounts for data-parallel sharding — each rank processes `1/nprocs` of the dataset.
|
||||
|
||||
> Document Update Time: 2026-07-30
|
||||
@@ -0,0 +1,235 @@
|
||||
# Getting Started
|
||||
|
||||
This guide walks you through installing AstrAI, downloading a model, running inference, preprocessing data, and launching your first training job.
|
||||
|
||||
## Prerequisites
|
||||
|
||||
- **Python 3.12+**
|
||||
- **PyTorch 2.11+** (CUDA 12.8 recommended for GPU support)
|
||||
- NVIDIA GPU with CUDA (optional but recommended; CPU works for inference)
|
||||
|
||||
## 1. Install
|
||||
|
||||
```bash
|
||||
git clone https://github.com/ViperEkura/AstrAI.git
|
||||
cd AstrAI
|
||||
|
||||
# Basic install (pure PyTorch, no custom CUDA kernels)
|
||||
pip install -e .
|
||||
|
||||
# With CUDA kernels (optional, for fused attention)
|
||||
# CSRC_KERNELS=true pip install -e . --no-build-isolation
|
||||
|
||||
# With dev dependencies (pytest, ruff)
|
||||
# pip install -e ".[dev]"
|
||||
```
|
||||
|
||||
> **CUDA kernels** are opt-in. They are not built by default. When built, they can be activated via `with attn_backend(ATTN_BACKEND.CUDA):` for accelerated decode/prefill. You can skip them for normal usage.
|
||||
|
||||
## 2. Download Model Weights
|
||||
|
||||
AstrAI uses HuggingFace-style model directories. Download the default 1B instruction-tuned model:
|
||||
|
||||
```bash
|
||||
python scripts/demo/download.py
|
||||
# → Downloads to params/
|
||||
```
|
||||
|
||||
To use a different model:
|
||||
|
||||
```bash
|
||||
python scripts/demo/download.py --repo-id <HF_REPO_ID> --local-dir ./my_model
|
||||
```
|
||||
|
||||
The model directory contains:
|
||||
- `config.json` — model architecture configuration
|
||||
- `model.safetensors` — model weights
|
||||
- `tokenizer.json` + `tokenizer_config.json` — tokenizer files (including chat template)
|
||||
|
||||
## 3. Run Inference
|
||||
|
||||
### Interactive Chat (Simplest)
|
||||
|
||||
```bash
|
||||
python scripts/demo/stream_chat.py
|
||||
# Type your message after >>, type !exit to quit
|
||||
```
|
||||
|
||||
This starts a multi-turn interactive chat session with streaming output.
|
||||
|
||||
### Start an HTTP Server
|
||||
|
||||
```bash
|
||||
# Terminal 1: start server
|
||||
python scripts/tools/server.py --param_path ./params --device cuda
|
||||
|
||||
# Terminal 2: query (OpenAI-compatible API)
|
||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
|
||||
```
|
||||
|
||||
The server also supports the Anthropic API at `/v1/messages`. See [Inference Guide](guides/inference.md) for full API documentation.
|
||||
|
||||
### Batch Generation from a File
|
||||
|
||||
Create an input JSONL file (one JSON object per line):
|
||||
|
||||
```json
|
||||
{"question": "What is machine learning?"}
|
||||
{"question": "Explain gradient descent."}
|
||||
```
|
||||
|
||||
```bash
|
||||
python scripts/tools/generate.py \
|
||||
--param_path ./params \
|
||||
--input_json_file input.jsonl \
|
||||
--output_json_file output.jsonl
|
||||
```
|
||||
|
||||
## 4. Preprocess Data
|
||||
|
||||
AstrAI uses a declarative JSON config to define the preprocessing pipeline. Create a config file for your training type:
|
||||
|
||||
### Pretraining (seq)
|
||||
|
||||
Input JSONL:
|
||||
```json
|
||||
{"text": "Artificial intelligence is..."}
|
||||
```
|
||||
|
||||
Config (`pretrain.json`):
|
||||
```json
|
||||
{
|
||||
"input": {
|
||||
"sections": [{"field": "text", "action": "train"}]
|
||||
},
|
||||
"preprocessing": {"max_seq_len": 2048},
|
||||
"output": {"storage_format": "bin"}
|
||||
}
|
||||
```
|
||||
|
||||
### SFT (Supervised Fine-Tuning)
|
||||
|
||||
Input JSONL:
|
||||
```json
|
||||
{"messages": [{"role": "user", "content": "Hi"}, {"role": "assistant", "content": "Hello!"}]}
|
||||
```
|
||||
|
||||
Config (`sft.json`):
|
||||
```json
|
||||
{
|
||||
"input": {
|
||||
"sections": [{"field": "messages", "action": "$role", "template": true}]
|
||||
},
|
||||
"mask": {
|
||||
"system": "mask",
|
||||
"user": "mask",
|
||||
"assistant": "train"
|
||||
},
|
||||
"mask_default": "mask",
|
||||
"preprocessing": {"max_seq_len": 2048},
|
||||
"output": {"storage_format": "bin", "dtype": {"loss_mask": "bool"}}
|
||||
}
|
||||
```
|
||||
|
||||
### Run Preprocessing
|
||||
|
||||
```bash
|
||||
python scripts/tools/preprocess.py data/*.jsonl -o output/ -c pretrain.json
|
||||
```
|
||||
|
||||
See [Preprocessing Guide](guides/preprocessing.md) for DPO/GRPO configs and all options.
|
||||
|
||||
## 5. Train
|
||||
|
||||
### Single GPU
|
||||
|
||||
```bash
|
||||
python scripts/tools/train.py \
|
||||
--train_type=seq \
|
||||
--data_root_path=/path/to/dataset \
|
||||
--param_path=./params \
|
||||
--batch_per_device=4 \
|
||||
--grad_accum_steps=8 \
|
||||
--max_lr=1e-4 \
|
||||
--window_size=2048 \
|
||||
--ckpt_dir=./checkpoint \
|
||||
--nprocs=1 \
|
||||
--parallel_mode=none
|
||||
```
|
||||
|
||||
### Multi-GPU (DDP)
|
||||
|
||||
```bash
|
||||
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export NCCL_NET_GDR_LEVEL=0
|
||||
|
||||
python scripts/tools/train.py \
|
||||
--train_type=seq \
|
||||
--data_root_path=/path/to/dataset \
|
||||
--param_path=./params \
|
||||
--parallel_mode=ddp \
|
||||
--nprocs=4 \
|
||||
--batch_per_device=4 \
|
||||
--grad_accum_steps=8 \
|
||||
--max_lr=1e-4 \
|
||||
--window_size=2048 \
|
||||
--ckpt_dir=./checkpoint
|
||||
```
|
||||
|
||||
### Training Types
|
||||
|
||||
| `--train_type` | Description | Data Keys |
|
||||
|----------------|-------------|-----------|
|
||||
| `seq` | Pre-training (next-token prediction) | `sequence` |
|
||||
| `sft` | Supervised fine-tuning (masked loss) | `sequence`, `loss_mask` |
|
||||
| `dpo` | Direct Preference Optimization | `chosen`, `rejected`, `*_mask` |
|
||||
| `grpo` | Group Relative Policy Optimization | `prompts`, `responses`, `masks`, `rewards` |
|
||||
|
||||
See [Training Guide](guides/training.md) for loss formulas and strategies. See [Distributed Guide](guides/distributed.md) for DDP/FSDP details.
|
||||
|
||||
## 6. Evaluate
|
||||
|
||||
```bash
|
||||
# HumanEval (code generation, auto-downloads dataset)
|
||||
python scripts/eval/evaluate_humaneval.py --param_path ./params --num_samples 20
|
||||
|
||||
# MMLU (knowledge, auto-downloads dataset)
|
||||
python scripts/eval/evaluate_mmlu.py --param_path ./params --n_shot 5
|
||||
|
||||
# Perplexity on custom data
|
||||
python scripts/eval/evaluate_ppl.py --param_path ./params --input_path data.jsonl --output_dir ppl_results/
|
||||
```
|
||||
|
||||
See [Evaluation Guide](guides/evaluation.md) for all benchmarks.
|
||||
|
||||
## 7. Docker
|
||||
|
||||
```bash
|
||||
# Build
|
||||
docker build -t astrai:latest .
|
||||
|
||||
# Run inference server with GPU
|
||||
docker run --gpus all -p 8000:8000 astrai:latest \
|
||||
python -m scripts.tools.server --port 8000 --device cuda
|
||||
|
||||
# Docker Compose (GPU)
|
||||
docker compose up -d
|
||||
```
|
||||
|
||||
## Next Steps
|
||||
|
||||
| Topic | Document |
|
||||
|-------|----------|
|
||||
| CLI parameters (train, server, generate, preprocess) | [CLI Reference](guides/params.md) |
|
||||
| Preprocessing pipeline details | [Preprocessing Guide](guides/preprocessing.md) |
|
||||
| Training loop, strategies, schedulers | [Training Guide](guides/training.md) |
|
||||
| KV cache, continuous batching, HTTP API | [Inference Guide](guides/inference.md) |
|
||||
| Evaluation benchmarks | [Evaluation Guide](guides/evaluation.md) |
|
||||
| Multi-GPU DDP / FSDP | [Distributed Guide](guides/distributed.md) |
|
||||
| System architecture | [Architecture](developer/architecture.md) |
|
||||
| Data pipeline internals | [Data Flow](developer/dataflow.md) |
|
||||
|
||||
> Document Update Time: 2026-07-30
|
||||
@@ -0,0 +1,255 @@
|
||||
# Distributed Training
|
||||
|
||||
AstrAI supports three parallel modes: **single GPU** (`none`), **Data Parallel** (`ddp`), and **Fully Sharded Data Parallel** (`fsdp`). This guide covers when to use each, how to launch multi-GPU training, and how gradient accumulation works.
|
||||
|
||||
## Quick Start
|
||||
|
||||
### Single GPU
|
||||
|
||||
```bash
|
||||
python scripts/tools/train.py \
|
||||
--train_type=sft \
|
||||
--param_path ./params \
|
||||
--data_root_path ./dataset \
|
||||
--parallel_mode=none \
|
||||
--nprocs=1 \
|
||||
--batch_per_device=4 \
|
||||
--grad_accum_steps=8
|
||||
```
|
||||
|
||||
### Multi-GPU DDP (4 GPUs)
|
||||
|
||||
```bash
|
||||
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export NCCL_NET_GDR_LEVEL=0
|
||||
|
||||
python scripts/tools/train.py \
|
||||
--train_type=sft \
|
||||
--param_path ./params \
|
||||
--data_root_path ./dataset \
|
||||
--parallel_mode=ddp \
|
||||
--nprocs=4 \
|
||||
--batch_per_device=4 \
|
||||
--grad_accum_steps=8
|
||||
```
|
||||
|
||||
### Multi-GPU FSDP (4 GPUs)
|
||||
|
||||
```bash
|
||||
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export NCCL_NET_GDR_LEVEL=0
|
||||
|
||||
python scripts/tools/train.py \
|
||||
--train_type=sft \
|
||||
--param_path ./params \
|
||||
--data_root_path ./dataset \
|
||||
--parallel_mode=fsdp \
|
||||
--nprocs=4 \
|
||||
--batch_per_device=4 \
|
||||
--grad_accum_steps=8
|
||||
```
|
||||
|
||||
> `--parallel_mode` defaults to `fsdp`. You can omit it for FSDP.
|
||||
|
||||
## Parallel Modes
|
||||
|
||||
| Mode | `--parallel_mode` | Param Layout | Memory | When to Use |
|
||||
|------|-------------------|--------------|--------|-------------|
|
||||
| Single GPU | `none` | Full, replicated | Highest | Small models, DPO/GRPO, debugging |
|
||||
| DDP | `ddp` | Full, replicated | High | Most multi-GPU training |
|
||||
| FSDP | `fsdp` | Sharded (DTensor) | Lowest | Large models that don't fit in single GPU |
|
||||
|
||||
### NoneExecutor
|
||||
|
||||
No wrapping. The model runs as-is on a single device. Gradient accumulation still works via `AccumOptimizer`/`AccumScheduler` (they gate `step()` on the sync counter). Checkpoint saving is a plain `state_dict()` call.
|
||||
|
||||
### DDPExecutor
|
||||
|
||||
Wraps the model with `torch.nn.parallel.DistributedDataParallel`. Each rank has a full copy of the model; gradients are all-reduced across ranks. Uses `gradient_as_bucket_view=True` and `broadcast_buffers=False` by default (hardcoded in `train.py`).
|
||||
|
||||
During gradient accumulation, non-sync micro-steps use `model.no_sync()` to skip gradient all-reduce. Only the final micro-step triggers the all-reduce.
|
||||
|
||||
### FSDPExecutor (FSDP2 / `fully_shard`)
|
||||
|
||||
Uses PyTorch's FSDP2 per-module API (`torch.distributed.fsdp.fully_shard`). Each model child (e.g., each `DecoderBlock`) is individually sharded — parameters become `DTensor`s distributed across ranks. No `FlatParameter`, original parameter names are preserved.
|
||||
|
||||
Key differences from DDP:
|
||||
- **Lower memory**: parameters are sharded, not replicated.
|
||||
- **Custom grad norm**: FSDP gradients are `DTensor`s, so `clip_grad_norm` computes the local norm, then all-reduces to get the global norm.
|
||||
- **Collective checkpoint ops**: `unshard()` and `full_tensor()` are collective — all ranks must call them even though only rank-0 saves. The executor handles this via `dist.barrier()` in `checkpoint_context`.
|
||||
- **Root skipped**: `fully_shard` is applied to direct children only (not the root model) due to an `ABC + Generic[T]` MRO incompatibility.
|
||||
|
||||
## Gradient Accumulation
|
||||
|
||||
Gradient accumulation lets you simulate a larger effective batch size by accumulating gradients over multiple micro-batches before calling `optimizer.step()`.
|
||||
|
||||
```
|
||||
Effective batch = nprocs × batch_per_device × grad_accum_steps
|
||||
```
|
||||
|
||||
Example: 4 GPUs × batch 4 × accum 8 = effective batch 256.
|
||||
|
||||
### How it works
|
||||
|
||||
Three cooperating layers:
|
||||
|
||||
1. **`GradientState`** — tracks the micro-step counter. Fires `sync_gradients=True` every `grad_accum_steps` micro-batches.
|
||||
2. **`executor._no_sync(model)`** — suppresses gradient synchronization on non-sync micro-steps:
|
||||
- `none`: `nullcontext` (nothing to skip)
|
||||
- `ddp`: `model.no_sync()` (skips all-reduce)
|
||||
- `fsdp`: `set_requires_gradient_sync(False)` on each `FSDPModule`
|
||||
3. **`AccumOptimizer` / `AccumScheduler`** — gate `step()` and `zero_grad()` on `sync_gradients`, so the optimizer only fires on the last micro-step.
|
||||
|
||||
The loss is divided by `grad_accum_steps` before `backward()`, so gradients sum to the correct mean.
|
||||
|
||||
## Process Launching
|
||||
|
||||
AstrAI auto-detects the launch method:
|
||||
|
||||
| Detection | Strategy | Use Case |
|
||||
|-----------|----------|----------|
|
||||
| `torchelastic` / `torchrun` env vars | `TorchrunStrategy` | External orchestrator (torchrun, SLURM, K8s) |
|
||||
| `RANK` + `WORLD_SIZE` env vars | `TorchrunStrategy` | External launch |
|
||||
| Neither | `LocalStrategy` | `python scripts/tools/train.py` (in-process spawn) |
|
||||
|
||||
### Local (default)
|
||||
|
||||
When you run `python scripts/tools/train.py --nprocs=4`, AstrAI uses `torch.multiprocessing.start_processes` to spawn 4 child processes. The parent process manages signal forwarding (SIGTERM/SIGINT) and waits for all children to finish.
|
||||
|
||||
### Torchrun
|
||||
|
||||
For multi-node or SLURM environments:
|
||||
|
||||
```bash
|
||||
torchrun --nproc_per_node=4 scripts/tools/train.py \
|
||||
--train_type=sft \
|
||||
--parallel_mode=ddp \
|
||||
--param_path ./params \
|
||||
--data_root_path ./dataset \
|
||||
--batch_per_device=4
|
||||
```
|
||||
|
||||
When launched via torchrun, AstrAI reads `RANK`, `WORLD_SIZE`, `LOCAL_RANK` from the environment and uses `TorchrunStrategy`. The `--nprocs` flag is ignored (the orchestrator controls process count).
|
||||
|
||||
## NCCL Environment Variables
|
||||
|
||||
For multi-GPU training, you **must** set these environment variables:
|
||||
|
||||
```bash
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export NCCL_NET_GDR_LEVEL=0
|
||||
```
|
||||
|
||||
These are required on certain hardware configurations (see `AGENTS.md`). Without them, NCCL may hang or crash during collective operations. These are set in the training shell scripts (`train-seq.sh`, `train-sft.sh`, `train-dpo.sh`) but not in Python code — you must export them before launching.
|
||||
|
||||
## Checkpoint Saving
|
||||
|
||||
Checkpoints are saved by **rank-0 only**. The flow:
|
||||
|
||||
1. `executor.checkpoint_context(model)` — wraps with `dist.barrier()` before and after (distributed only).
|
||||
2. `executor.unwrap_model(model)` — gathers the full state dict:
|
||||
- `none`: `model.state_dict()`
|
||||
- `ddp`: `model.module.state_dict()`
|
||||
- `fsdp`: `unshard()` → `full_tensor()` → `reshard()` (collective on all ranks, result kept only on rank-0)
|
||||
3. Non-rank-0 ranks get `None` — the save is skipped.
|
||||
4. Rank-0 writes `meta.json`, `config.json`, `model.safetensors`, and optional `{key}.pt` (optimizer/scheduler state).
|
||||
|
||||
> **FSDP note**: Even though only rank-0 saves, all ranks must participate in `unwrap_model` because `unshard()` and `full_tensor()` are collective operations. The barriers in `checkpoint_context` keep all ranks in lockstep.
|
||||
|
||||
## Total Steps Calculation
|
||||
|
||||
The scheduler's total step count accounts for data-parallel sharding:
|
||||
|
||||
```
|
||||
samples_per_replica = ceil(dataset_len / nprocs)
|
||||
batches_per_replica = ceil(samples_per_replica / batch_per_device)
|
||||
total_steps = (batches_per_replica // grad_accum_steps) * n_epoch
|
||||
```
|
||||
|
||||
This ensures the LR schedule is correctly scaled regardless of the number of GPUs.
|
||||
|
||||
## Real Examples
|
||||
|
||||
### Pretraining (seq, DDP, 4 GPUs)
|
||||
|
||||
```bash
|
||||
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export NCCL_NET_GDR_LEVEL=0
|
||||
|
||||
python scripts/tools/train.py \
|
||||
--train_type=seq \
|
||||
--param_path ./params \
|
||||
--data_root_path ./dataset/cached \
|
||||
--parallel_mode=ddp \
|
||||
--nprocs=4 \
|
||||
--n_epoch=1 \
|
||||
--max_lr=2e-4 \
|
||||
--schedule_type=wsd \
|
||||
--warmup_ratio=0.02 \
|
||||
--window_size=2048 \
|
||||
--batch_per_device=4 \
|
||||
--grad_accum_steps=32 \
|
||||
--ckpt_interval=2000
|
||||
# Effective batch = 4 × 4 × 32 = 512
|
||||
```
|
||||
|
||||
### SFT (DDP, 4 GPUs)
|
||||
|
||||
```bash
|
||||
python scripts/tools/train.py \
|
||||
--train_type=sft \
|
||||
--param_path ./AstrAI-V1-base \
|
||||
--data_root_path ./dataset/cached_sft \
|
||||
--parallel_mode=ddp \
|
||||
--nprocs=4 \
|
||||
--n_epoch=2 \
|
||||
--max_lr=2e-5 \
|
||||
--schedule_type=cosine \
|
||||
--warmup_ratio=0.02 \
|
||||
--min_rate=0.05 \
|
||||
--window_size=2048 \
|
||||
--batch_per_device=4 \
|
||||
--grad_accum_steps=8
|
||||
# Effective batch = 4 × 4 × 8 = 128
|
||||
```
|
||||
|
||||
### DPO (Single GPU)
|
||||
|
||||
```bash
|
||||
python scripts/tools/train.py \
|
||||
--train_type=dpo \
|
||||
--param_path ./checkpoint/epoch_1_step_6000 \
|
||||
--data_root_path ./alpaca_dpo.jsonl \
|
||||
--parallel_mode=none \
|
||||
--nprocs=1 \
|
||||
--max_lr=5e-6 \
|
||||
--schedule_type=cosine \
|
||||
--warmup_ratio=0.1 \
|
||||
--min_rate=0.1 \
|
||||
--window_size=1024 \
|
||||
--batch_per_device=4 \
|
||||
--grad_accum_steps=8 \
|
||||
--dpo_beta=0.1 \
|
||||
--max_grad_norm=50
|
||||
```
|
||||
|
||||
## CLI Parameters
|
||||
|
||||
| Parameter | Default | Description |
|
||||
|-----------|---------|-------------|
|
||||
| `--nprocs` | 1 | Number of GPUs / processes |
|
||||
| `--parallel_mode` | `fsdp` | `none`, `ddp`, or `fsdp` |
|
||||
| `--start_method` | `spawn` | Multiprocessing start method (`spawn`, `fork`, `forkserver`) |
|
||||
| `--backend` | `nccl` | Distributed backend (`nccl`, `gloo`) |
|
||||
| `--master_addr` | `localhost` | Master node address |
|
||||
| `--master_port` | `29500` | Master node port |
|
||||
| `--device_type` | `cuda` | Device type |
|
||||
|
||||
> `--tp_size` is parsed but **not yet wired** — tensor parallelism is future work. `ColumnParallelLinear` / `RowParallelLinear` exist in `astrai/parallel/module.py` but are not used by the model.
|
||||
|
||||
Full parameter reference: [CLI Reference](params.md). Training loop and strategies: [Training Guide](training.md).
|
||||
|
||||
> Document Update Time: 2026-07-30
|
||||
@@ -0,0 +1,252 @@
|
||||
# Evaluation
|
||||
|
||||
AstrAI provides 7 evaluation scripts in `scripts/eval/` covering code generation, knowledge QA, perplexity, summarization, data quality, instruction following, and weight analysis.
|
||||
|
||||
## Overview
|
||||
|
||||
| Script | Metric | Model Invocation | External Dataset |
|
||||
|--------|--------|-------------------|-------------------|
|
||||
| `evaluate_humaneval.py` | Code-gen pass@1/10/100 | `InferenceEngine.generate` | HF `openai/openai_humaneval` (auto-download) |
|
||||
| `evaluate_mmlu.py` | MCQ accuracy (log-likelihood) | Direct `model()` forward | HF `cais/mmlu` (auto-download) |
|
||||
| `evaluate_ppl.py` | Perplexity / token loss | Direct `model()` forward | User JSONL |
|
||||
| `evaluate_rouge.py` | ROUGE-1/2/L | None (pure metric) | User JSONL |
|
||||
| `evaluate_ifd.py` | Instruction-Following Difficulty | Direct `model()` forward | User JSONL |
|
||||
| `evaluate_ifeval.py` | Instruction-following constraints | `InferenceEngine.generate` | HF `google/IFEval` (auto-download) |
|
||||
| `analyze_weights.py` | SVD effective rank / weight stats | None (loads safetensors) | Checkpoint dir |
|
||||
|
||||
Two invocation patterns exist:
|
||||
- **Generation benchmarks** (HumanEval, IFEval): use `InferenceEngine` to generate responses, then score them.
|
||||
- **Scoring benchmarks** (MMLU, PPL, IFD): call `model()` directly under `torch.inference_mode()` for log-likelihood computation.
|
||||
|
||||
Common defaults: `--param_path` defaults to `./params`; dtype defaults to `bfloat16` on CUDA, `float32` on CPU.
|
||||
|
||||
---
|
||||
|
||||
## HumanEval (Code Generation)
|
||||
|
||||
Generates completions for 164 programming problems, executes them against hidden tests, and reports pass@k.
|
||||
|
||||
```bash
|
||||
python scripts/eval/evaluate_humaneval.py \
|
||||
--param_path ./params \
|
||||
--num_samples 20 \
|
||||
--batch_size 32 \
|
||||
--max_tokens 512 \
|
||||
--output results/humaneval.json
|
||||
```
|
||||
|
||||
| Parameter | Default | Description |
|
||||
|-----------|---------|-------------|
|
||||
| `--param_path` | `./params` | Model directory |
|
||||
| `--data_path` | `./humaneval/HumanEval.jsonl` | HumanEval JSONL (auto-downloaded if missing) |
|
||||
| `--output` | None | Save results JSON (also writes `_completions.json`) |
|
||||
| `--test_only` | None | Test an existing completions JSON (skip generation) |
|
||||
| `--generate_only` | False | Only generate, skip execution/testing |
|
||||
| `--num_samples` | 200 | Completions per problem (pass@k needs >= k) |
|
||||
| `--max_tokens` | 512 | Max generation length |
|
||||
| `--temperature` | 0.8 | Sampling temperature |
|
||||
| `--top_p` | 0.95 | Nucleus sampling threshold |
|
||||
| `--top_k` | 50 | Top-k sampling |
|
||||
| `--batch_size` | 32 | Generation batch size |
|
||||
| `--test_workers` | 8 | ProcessPoolExecutor workers for test execution |
|
||||
| `--test_timeout` | 3.0 | Per-subprocess timeout (seconds) |
|
||||
| `--problems` | None | Restrict to specific problem indices |
|
||||
|
||||
**Output**: stdout prints `pass@1`, `pass@10`, `pass@100`. With `--output`, writes per-problem results + `_summary` aggregate and a `_completions.json` file.
|
||||
|
||||
**Data**: Auto-downloads `openai/openai_humaneval` from HuggingFace on first run. Each problem has `task_id`, `entry_point`, `prompt`, `test`.
|
||||
|
||||
---
|
||||
|
||||
## MMLU (Knowledge QA)
|
||||
|
||||
57-subject multiple-choice accuracy via log-likelihood comparison. Supports n-shot few-shot prompting and option permutation.
|
||||
|
||||
```bash
|
||||
python scripts/eval/evaluate_mmlu.py \
|
||||
--param_path ./params \
|
||||
--n_shot 5 \
|
||||
--subjects math_algebra history_us \
|
||||
--output results/mmlu.json
|
||||
```
|
||||
|
||||
| Parameter | Default | Description |
|
||||
|-----------|---------|-------------|
|
||||
| `--param_path` | `./params` | Model directory |
|
||||
| `--data_dir` | `./mmlu_data` | MMLU data directory (per-subject CSVs) |
|
||||
| `--download` | False | Force re-download |
|
||||
| `--n_shot` | 5 | Few-shot examples (0 = zero-shot) |
|
||||
| `--subjects` | all 57 | Specific subjects to evaluate |
|
||||
| `--output` | None | Output JSON path |
|
||||
| `--split` | `test` | `test` or `val` |
|
||||
| `--device` | auto | Device (`cuda` / `cpu`) |
|
||||
| `--dtype` | auto | `bfloat16` on CUDA, `float32` on CPU |
|
||||
| `--seed` | 0 | Seed for option permutation (0 = enabled, -1 = disabled) |
|
||||
|
||||
**How it works**: For each question, builds a prompt with n-shot examples, then scores each choice (A/B/C/D) by computing the summed log-likelihood of the choice token given the context. The choice with the highest log-prob is the prediction.
|
||||
|
||||
**Output**: stdout prints per-subject accuracy and overall. With `--output`, writes per-subject `{accuracy, correct, total}` + `_overall` aggregate.
|
||||
|
||||
**Data**: Auto-downloads `cais/mmlu` from HuggingFace. Stored as per-subject CSVs in `<data_dir>/<split>/` and `<data_dir>/dev/` (for few-shot).
|
||||
|
||||
---
|
||||
|
||||
## Perplexity (PPL)
|
||||
|
||||
Token-level negative-log-likelihood and perplexity on arbitrary text data. Supports streaming mode (memory-efficient) and non-streaming mode (exact per-token stats).
|
||||
|
||||
```bash
|
||||
python scripts/eval/evaluate_ppl.py \
|
||||
--param_path ./params \
|
||||
--input_path data.jsonl \
|
||||
--output_dir ppl_results/ \
|
||||
--batch_size 4 \
|
||||
--max_length 2048
|
||||
```
|
||||
|
||||
| Parameter | Default | Description |
|
||||
|-----------|---------|-------------|
|
||||
| `--param_path` | required | Model directory |
|
||||
| `--input_path` | required | Input file, glob, or directory |
|
||||
| `--output_dir` | required | Output directory for `summary.json` + token JSONL |
|
||||
| `--text_key` | `text` | Key for the text field in input data |
|
||||
| `--batch_size` | 4 | Batch size |
|
||||
| `--max_length` | 2048 | Max sequence length (tokens) |
|
||||
| `--token_level` | False | Store per-token log_probs + token-type analysis |
|
||||
| `--max_samples` | None | Random subsample per file |
|
||||
| `--device` | auto | Device |
|
||||
| `--dtype` | auto | Torch dtype |
|
||||
|
||||
**Input**: JSONL or JSON files. Each item must have a field named by `--text_key` (default `text`). If `--input_path` is a directory, recursively collects `*.jsonl` and `*.json`.
|
||||
|
||||
**Output**: `summary.json` with per-file stats (tokens, mean/median loss, perplexity, p50/p90/p95/p99). With `--token_level`, also writes per-token JSONL with token IDs and log-probs.
|
||||
|
||||
---
|
||||
|
||||
## ROUGE
|
||||
|
||||
ROUGE-1/2/L (precision, recall, F1) for summarization. Self-contained implementation with no external dependencies.
|
||||
|
||||
```bash
|
||||
python scripts/eval/evaluate_rouge.py \
|
||||
--data_path predictions.jsonl \
|
||||
--output results/rouge.json
|
||||
```
|
||||
|
||||
| Parameter | Default | Description |
|
||||
|-----------|---------|-------------|
|
||||
| `--data_path` | required | JSONL with `reference`/`candidate` per line |
|
||||
| `--output` | None | Output JSON path |
|
||||
|
||||
**Input**: JSONL, one object per line:
|
||||
```json
|
||||
{"reference": "Ground truth text", "candidate": "Model output text"}
|
||||
```
|
||||
|
||||
**Output**: stdout prints `rouge-1`, `rouge-2`, `rouge-l` each as P/R/F1. With `--output`, writes JSON with `aggregate` and `per_item` scores.
|
||||
|
||||
Can also be imported as a library:
|
||||
```python
|
||||
from scripts.eval.evaluate_rouge import compute_rouge
|
||||
scores = compute_rouge(reference, candidate)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## IFD (Instruction-Following Difficulty)
|
||||
|
||||
Data quality metric: `IFD = L_conditional / L_unconditional`. Measures how much harder it is to predict a response given its instruction vs. without it. Useful for filtering instruction-tuning data.
|
||||
|
||||
```bash
|
||||
python scripts/eval/evaluate_ifd.py \
|
||||
--param_path ./params \
|
||||
--input_path sft_data.jsonl \
|
||||
--output_dir ifd_results/ \
|
||||
--format messages \
|
||||
--batch_size 8
|
||||
```
|
||||
|
||||
| Parameter | Default | Description |
|
||||
|-----------|---------|-------------|
|
||||
| `--param_path` | required | Model directory |
|
||||
| `--input_path` | required | Input file, glob, or directory |
|
||||
| `--output_dir` | required | Output directory |
|
||||
| `--max_len` | 2048 | Max token length |
|
||||
| `--format` | `plain` | `plain` (instruction/response fields) or `messages` (chat format) |
|
||||
| `--instr_key` | `instruction` | Instruction field key (plain format) |
|
||||
| `--resp_key` | `response` | Response field key (plain format) |
|
||||
| `--batch_size` | 8 | Items per model-forward flush |
|
||||
| `--device` | auto | Device |
|
||||
| `--dtype` | auto | Torch dtype |
|
||||
| `--sentinel_text` | `\n` | Prefix for unconditional pass (`""` → bos/pad fallback) |
|
||||
| `--per_token` | False | Include per-token IFD breakdown |
|
||||
| `--max_samples` | None | Random subsample per file |
|
||||
|
||||
**How it works**: Two forward passes per batch — (1) conditional: packed BFD sequence with context + response, (2) unconditional: response prefixed with a sentinel. IFD = mean_conditional_loss / mean_unconditional_loss. IFD > 1 means the instruction makes the response harder to predict (higher quality data).
|
||||
|
||||
**Output**: Per-file `<label>_ifd.jsonl` with IFD scores per item. `summary.json` aggregates per-file stats.
|
||||
|
||||
---
|
||||
|
||||
## IFEval (Instruction Following)
|
||||
|
||||
Google's IFEval benchmark: generates responses and verifies 27 types of constraints (keywords, format, length, case, punctuation, etc.).
|
||||
|
||||
```bash
|
||||
python scripts/eval/evaluate_ifeval.py \
|
||||
--param_path ./params \
|
||||
--num_samples 1 \
|
||||
--max_tokens 512 \
|
||||
--output results/ifeval.json
|
||||
```
|
||||
|
||||
| Parameter | Default | Description |
|
||||
|-----------|---------|-------------|
|
||||
| `--param_path` | `./params` | Model directory |
|
||||
| `--data_path` | `./ifeval/input_data.jsonl` | IFEval JSONL (auto-downloaded if missing) |
|
||||
| `--output` | None | Output JSON path |
|
||||
| `--max_tokens` | 512 | Max generation tokens |
|
||||
| `--temperature` | 0.1 | Sampling temperature (low for instruction-following) |
|
||||
| `--top_p` | 0.95 | Top-p sampling |
|
||||
| `--top_k` | 50 | Top-k sampling |
|
||||
| `--num_samples` | 1 | Samples per problem (best-of-n scoring) |
|
||||
| `--batch_size` | 1 | Inference batch size |
|
||||
| `--limit` | None | Limit to first N problems (quick testing) |
|
||||
| `--dump_responses` | None | Path to dump raw responses as JSONL |
|
||||
|
||||
**Output**: stdout prints overall accuracy + per-constraint-type accuracy table. With `--output`, writes per-problem results + `_summary`.
|
||||
|
||||
**Data**: Auto-downloads `google/IFEval` from HuggingFace. Each problem has `key`, `prompt`, `instruction_id_list`, `kwargs`.
|
||||
|
||||
---
|
||||
|
||||
## Weight Analysis
|
||||
|
||||
SVD-based effective rank and weight statistics for checkpoint diagnostics. Does not load the model graph or run any forward pass.
|
||||
|
||||
```bash
|
||||
python scripts/eval/analyze_weights.py \
|
||||
--ckpt_dir ./checkpoint/epoch_1_step_6000 \
|
||||
--output results/weights.json
|
||||
```
|
||||
|
||||
| Parameter | Default | Description |
|
||||
|-----------|---------|-------------|
|
||||
| `--ckpt_dir` | required | Checkpoint dir with `model.safetensors` + `config.json` |
|
||||
| `--compare` | None | Additional checkpoint dirs to compare |
|
||||
| `--no_svd` | False | Skip SVD; show only weight stats (faster) |
|
||||
| `--output` | None | Save results as JSON |
|
||||
| `--device` | `cuda` | Device for SVD |
|
||||
|
||||
**Output**: SVD effective rank by component (ER@90/95/99%, entropic rank, condition number), per-layer effective rank grid, and weight value statistics (mean/std/min/max). Provides a utilization verdict (HIGH >0.85 / MODERATE >0.5 / LOW).
|
||||
|
||||
---
|
||||
|
||||
## Tips
|
||||
|
||||
- **Quick test**: Use `--limit` (IFEval) or `--problems` (HumanEval) to run on a small subset first.
|
||||
- **Auto-download**: HumanEval, MMLU, and IFEval auto-download their datasets on first run. The other scripts expect user-provided data.
|
||||
- **Output formats**: `--output` writes a single JSON for most scripts. PPL and IFD write an `--output_dir` containing `summary.json` plus per-file artifacts.
|
||||
- **CPU mode**: All scripts auto-detect CUDA. To force CPU, use `--device cpu --dtype float32`.
|
||||
|
||||
> Document Update Time: 2026-07-30
|
||||
@@ -4,6 +4,7 @@
|
||||
|
||||
- [KV Cache](#kv-cache)
|
||||
- [KVCache System](#kvcache-system)
|
||||
- [Attention Backend](#attention-backend)
|
||||
- [Continuous Batching](#continuous-batching)
|
||||
- [Sampling](#sampling-strategy-pattern)
|
||||
- [Protocol Handlers](#protocol-handlers-strategy-pattern)
|
||||
@@ -23,30 +24,58 @@ RoPE is applied **before** KV cache write, not after — otherwise position enco
|
||||
|
||||
## KVCache System
|
||||
|
||||
Seven classes working together, with two concrete cache implementations:
|
||||
|
||||
### ContiguousCache (default)
|
||||
Three-layer separation (SGLang-inspired): storage, index table, allocator.
|
||||
|
||||
```
|
||||
ContiguousCache (simple contiguous per-slot cache)
|
||||
├── ContiguousCacheView bundles k/v tensors + slot indices for attention layers
|
||||
PagePool (top-level manager, orchestrates all layers)
|
||||
├── KVStorage k_buffer / v_buffer [n_layers, size, n_kv_heads, head_dim]
|
||||
├── ReqToTokenPool req_to_token [num_reqs, max_ctx_len] → physical token slot
|
||||
├── Allocator bitmask-based page allocator + ref-count + LRU (paged mode only)
|
||||
└── PrefixCache hash-based prefix matching (paged mode only)
|
||||
```
|
||||
|
||||
Created by default when no cache is passed to `InferenceScheduler`. Each task occupies a fixed slot of `[max_seq_len, num_key_value_heads, head_dim]`. Simple and efficient for small-to-medium batch sizes.
|
||||
`PagePool` supports two modes:
|
||||
|
||||
### PageCache (paged with prefix sharing)
|
||||
- **Contiguous (default)**: pre-allocates `max_batch_size * max_seq_len` token slots. `req_to_token` is a trivial linear mapping (`slot = req_idx * max_seq_len + pos`). No dynamic allocation.
|
||||
- **Paged** (`page_size=1` or `>1` with `n_tokens` set): shared token pool with on-demand allocation. Allocator + PrefixCache enable prefix sharing and LRU eviction.
|
||||
|
||||
`bind_tasks()` returns a `KVCache` dataclass — pure data, no methods:
|
||||
|
||||
```
|
||||
PageCache (paged KV cache with prefix sharing, alternative)
|
||||
├── PagePool orchestrates page allocation + prefix matching
|
||||
│ ├── Allocator bitmask-based page allocator + ref-count + LRU
|
||||
│ └── PrefixCache hash-based prefix matching (page_hash via polynomial hash)
|
||||
├── TaskTable maps task_id → page_table + cached token count
|
||||
├── Storage k_cache / v_cache tensors (num_hidden_layers × n_pages × page_size × num_key_value_heads × head_dim)
|
||||
└── PageCacheView bundles Storage + page_table + total_len for attention layers
|
||||
KVCache
|
||||
├── k_buffer, v_buffer [n_layers, size, n_kv_heads, head_dim]
|
||||
├── req_to_token [num_reqs, max_ctx_len]
|
||||
├── req_pool_indices [batch_size]
|
||||
├── seq_lens [batch_size]
|
||||
└── out_cache_loc [batch, seq_len] — write indices for this forward
|
||||
```
|
||||
|
||||
`isinstance(cache, KVCache)` checks dispatch to the correct view. Both implement the abstract `KVCache` interface used by `Executor` and `InferenceScheduler`.
|
||||
Attention layers do raw buffer indexing: `k_buffer[layer_id, out_cache_loc] = k` to write, `k_buffer[layer_id, indices]` to gather.
|
||||
|
||||
## Attention Backend
|
||||
|
||||
Attention computation (cache I/O + SDPA/kernel dispatch) is decoupled from the model via `AttentionBackend` ABC:
|
||||
|
||||
```
|
||||
AttentionBackend (ABC)
|
||||
├── TorchNativeBackend SDPA + indirect KV cache gather (default)
|
||||
└── CudaBackend CUDA kernel dispatch (attn_paged_decode, attn_prefill)
|
||||
```
|
||||
|
||||
Select via context manager (mirrors `torch.nn.attention.sdpa_kernel`):
|
||||
|
||||
```python
|
||||
from astrai.extension import attn_backend, ATTN_BACKEND
|
||||
|
||||
with attn_backend(ATTN_BACKEND.CUDA):
|
||||
engine.generate("hello")
|
||||
```
|
||||
|
||||
`CudaBackend` decode path: writes K/V to cache, then calls `attn_paged_decode` with `page_size=1` — the `req_to_token` table serves directly as the page table, each token slot is a single-token "page". No explicit K/V gather needed.
|
||||
|
||||
`CudaBackend` prefill path: writes K/V, gathers full-sequence K/V via indirect indexing (same as `TorchNativeBackend`), then calls `attn_prefill`.
|
||||
|
||||
Fallback: `CudaBackend` delegates to `TorchNativeBackend` when a CUDA kernel is not available.
|
||||
|
||||
## Continuous Batching
|
||||
|
||||
@@ -249,4 +278,4 @@ async for token in engine.generate_async("Hello", ...): # -> AsyncGenerator[s
|
||||
print(token)
|
||||
```
|
||||
|
||||
> Document Update Time: 2026-07-09
|
||||
> Document Update Time: 2026-07-30
|
||||
@@ -84,7 +84,7 @@ Combined optimizer: matrix parameters via **Muon**, non-matrix via **AdamW** (`f
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--nprocs` | Number of GPUs / processes | 1 |
|
||||
| `--parallel_mode` | Parallel strategy (`none`, `ddp`, or `fsdp`) | none |
|
||||
| `--parallel_mode` | Parallel strategy (`none`, `ddp`, `fsdp`) | fsdp |
|
||||
| `--device_type` | Device type | cuda |
|
||||
| `--start_method` | Multiprocessing start method (`spawn`, `fork`, `forkserver`) | spawn |
|
||||
| `--backend` | Distributed training backend | nccl |
|
||||
@@ -164,6 +164,7 @@ nohup python scripts/tools/train.py \
|
||||
| `--device` | str | `cuda` | Device to load model on |
|
||||
| `--dtype` | str | `bfloat16` | Model weights dtype (`bfloat16`, `float16`, `float32`) |
|
||||
| `--max_batch_size` | int | `16` | Maximum batch size for continuous batching |
|
||||
| `--max_seq_len` | int | model config `max_position_embeddings` | Maximum sequence length (KV cache size + prompt truncation) |
|
||||
| `--reload` | flag | `False` | Enable auto-reload for development |
|
||||
|
||||
Usage:
|
||||
@@ -173,6 +174,14 @@ python scripts/tools/server.py --param_path ./params --device cuda --dtype bfloa
|
||||
|
||||
See [Inference Guide](inference.md) for HTTP API documentation.
|
||||
|
||||
# Preprocess
|
||||
|
||||
```bash
|
||||
python scripts/tools/preprocess.py data/*.jsonl -o output/ -c config.json
|
||||
```
|
||||
|
||||
See [Preprocessing Guide](preprocessing.md) for config file format and examples.
|
||||
|
||||
## Generate (`generate.py`)
|
||||
|
||||
| Parameter | Type | Default | Description |
|
||||
@@ -186,7 +195,11 @@ See [Inference Guide](inference.md) for HTTP API documentation.
|
||||
| `--top_k` | int | `30` | Top-k filtering |
|
||||
| `--top_p` | float | `0.95` | Nucleus sampling threshold |
|
||||
| `--batch_size` | int | `1` | Batch size for generation |
|
||||
| `--num_samples` | int | `1` | Responses per prompt |
|
||||
| `--max_tokens` | int | model config `max_position_embeddings` | Maximum tokens to generate |
|
||||
| `--cache_len` | int | `2048` | KV cache length |
|
||||
| `--frequency_penalty` | float | `0.0` | Frequency penalty |
|
||||
| `--rep_window` | int | `64` | Window size for frequency penalty |
|
||||
|
||||
Usage:
|
||||
```bash
|
||||
@@ -256,6 +256,7 @@ When `sources` is set, `sections` is ignored.
|
||||
| `min_chars` | int | `50` | Skip text-mode items shorter than this |
|
||||
| `max_chars` | int | `2000000` | Skip text-mode items longer than this |
|
||||
| `max_items` | int or null | `null` | Stop after N documents |
|
||||
| `batch_size` | int | `256` | Records per tokenization batch |
|
||||
| `packing_strategy` | str | `"simple"` | Packing strategy: `"simple"`, `"bfd"`, `"bfd_split"` |
|
||||
| `max_packed_len` | int | `8192` | Maximum length of a packed bin |
|
||||
| `truncation_mode` | str | `"keep_start"` | How to truncate sequences: `"keep_start"` or `"keep_end"` |
|
||||
@@ -265,7 +266,7 @@ When `sources` is set, `sections` is ignored.
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `domain_key` | str or null | `null` | JSONL key for domain grouping |
|
||||
| `storage_format` | str | `"bin"` | `"bin"` (mmap) or `"h5"` |
|
||||
| `storage_format` | str | `"bin"` | `"bin"` (mmap). Reading also supports `"jsonl"` for on-the-fly tokenization |
|
||||
| `max_tokens_per_shard` | int | `100000000` | Flush threshold in cumulative tokens |
|
||||
| `dtype` | dict[str, str] | `{}` | Per-key tensor dtype override (e.g. `{"loss_mask": "bool"}`) |
|
||||
| `position_ids_mode` | str | `"doc_reset"` | How to compute position_ids: `"none"`, `"doc_reset"`, `"continuous"` |
|
||||
|
Before Width: | Height: | Size: 281 KiB After Width: | Height: | Size: 281 KiB |
+4
-4
@@ -8,7 +8,6 @@ name = "astrai"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.12"
|
||||
dependencies = [
|
||||
"h5py==3.15.1",
|
||||
"numpy==2.4.4",
|
||||
"torch==2.11.0",
|
||||
"tokenizers==0.21.4",
|
||||
@@ -16,10 +15,11 @@ dependencies = [
|
||||
"safetensors==0.5.3",
|
||||
"huggingface-hub==0.34.3",
|
||||
"jinja2>=3.0.0",
|
||||
"pydantic>=2.0",
|
||||
"fastapi",
|
||||
"uvicorn[standard]",
|
||||
"httpx",
|
||||
"requests",
|
||||
"click>=8.0",
|
||||
"pyyaml>=6.0",
|
||||
]
|
||||
keywords = ["nlp", "datasets", "language-models", "machine-learning"]
|
||||
license = { text = "GPL-3.0" }
|
||||
@@ -31,7 +31,7 @@ classifiers = [
|
||||
urls = { Homepage = "https://github.com/ViperEkura/AstrAI" }
|
||||
|
||||
[project.optional-dependencies]
|
||||
dev = ["pytest==9.0.2", "ruff"]
|
||||
dev = ["pytest==9.0.2", "ruff", "httpx2"]
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
where = ["."]
|
||||
|
||||
@@ -8,11 +8,11 @@ import safetensors.torch
|
||||
import torch
|
||||
|
||||
|
||||
def effective_rank_metrics(w: torch.Tensor) -> dict:
|
||||
def effective_rank_metrics(w: torch.Tensor, device: str = "cpu") -> dict:
|
||||
if w.ndim == 1:
|
||||
return {"shape": tuple(w.shape), "is_1d": True}
|
||||
|
||||
w = w.float()
|
||||
w = w.float().to(device)
|
||||
s = torch.linalg.svdvals(w)
|
||||
s_sq = s**2
|
||||
total = s_sq.sum()
|
||||
@@ -238,6 +238,12 @@ def main():
|
||||
default=None,
|
||||
help="Save results as JSON to this path.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--device",
|
||||
type=str,
|
||||
default="cuda",
|
||||
help="Device for SVD computation (e.g., 'cuda:0', 'cpu').",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
all_results = {}
|
||||
@@ -277,10 +283,12 @@ def main():
|
||||
|
||||
results = {}
|
||||
if not args.no_svd:
|
||||
print(f"Computing SVD on {len(weight_keys)} tensors...")
|
||||
print(
|
||||
f"Computing SVD on {len(weight_keys)} tensors (device={args.device})..."
|
||||
)
|
||||
for i, k in enumerate(sorted(weight_keys)):
|
||||
print(f" [{i + 1}/{len(weight_keys)}] {k:<60s}", end="\r")
|
||||
results[k] = effective_rank_metrics(sd[k])
|
||||
results[k] = effective_rank_metrics(sd[k], device=args.device)
|
||||
print()
|
||||
else:
|
||||
print(f"Computing stats on {len(weight_keys)} tensors (no SVD)...")
|
||||
|
||||
@@ -57,6 +57,7 @@ class EvalConfig:
|
||||
top_p: float = 0.95
|
||||
top_k: int = 50
|
||||
batch_size: int = 32
|
||||
max_seq_len: int = 4096
|
||||
test_timeout: float = 3.0
|
||||
test_workers: int = 8
|
||||
k_values: Tuple[int, ...] = (1, 10, 100)
|
||||
@@ -90,7 +91,9 @@ def save_json(path: str, data):
|
||||
json.dump(data, f, indent=2, ensure_ascii=False)
|
||||
|
||||
|
||||
def create_engine(param_path: str, batch_size: int) -> InferenceEngine:
|
||||
def create_engine(
|
||||
param_path: str, batch_size: int, max_seq_len: int
|
||||
) -> InferenceEngine:
|
||||
model = AutoModel.from_pretrained(param_path)
|
||||
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
||||
model.to(device="cuda", dtype=torch.bfloat16)
|
||||
@@ -98,6 +101,7 @@ def create_engine(param_path: str, batch_size: int) -> InferenceEngine:
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=batch_size,
|
||||
max_seq_len=max_seq_len,
|
||||
)
|
||||
|
||||
|
||||
@@ -318,7 +322,7 @@ def run_pipeline(cfg: EvalConfig) -> Dict:
|
||||
if cfg.problem_indices:
|
||||
problems = [problems[i] for i in cfg.problem_indices if i < len(problems)]
|
||||
|
||||
engine = create_engine(cfg.param_path, cfg.batch_size)
|
||||
engine = create_engine(cfg.param_path, cfg.batch_size, cfg.max_seq_len)
|
||||
|
||||
try:
|
||||
generated = generate_all(engine, problems, cfg)
|
||||
@@ -357,7 +361,8 @@ def parse_args(argv: Optional[List[str]] = None) -> EvalConfig:
|
||||
p.add_argument("--temperature", type=float, default=0.8)
|
||||
p.add_argument("--top_p", type=float, default=0.95)
|
||||
p.add_argument("--top_k", type=int, default=50)
|
||||
p.add_argument("--batch_size", type=int, default=32)
|
||||
p.add_argument("--batch_size", type=int, default=64)
|
||||
p.add_argument("--max_seq_len", type=int, default=4096)
|
||||
p.add_argument("--test_workers", type=int, default=8)
|
||||
p.add_argument("--test_timeout", type=float, default=3.0)
|
||||
p.add_argument("--problems", type=int, nargs="+", default=None)
|
||||
@@ -375,6 +380,7 @@ def parse_args(argv: Optional[List[str]] = None) -> EvalConfig:
|
||||
top_p=args.top_p,
|
||||
top_k=args.top_k,
|
||||
batch_size=args.batch_size,
|
||||
max_seq_len=args.max_seq_len,
|
||||
test_workers=args.test_workers,
|
||||
test_timeout=args.test_timeout,
|
||||
problem_indices=args.problems,
|
||||
|
||||
@@ -9,10 +9,16 @@ v2 changelog:
|
||||
- Same token set: unconditional pass prefixes resp with a plain-text sentinel
|
||||
(default ``\\n``; use ``--sentinel_text ""`` for bos/pad fallback).
|
||||
Both branches predict the identical N resp tokens.
|
||||
Single-token answers (rl=1) are now supported.
|
||||
- Single-token answers (rl=1) are now supported.
|
||||
- ctx_len tracked in output
|
||||
- skip_reason for None samples (no more silent None)
|
||||
- --per_token for per-token IFD breakdown
|
||||
|
||||
v3 changelog:
|
||||
- Append EOS at the end of response in both conditional and unconditional
|
||||
passes (``--append_eos`` / ``--no-append_eos``, default: enabled).
|
||||
The model now also predicts when the response should end, which is part
|
||||
of instruction following. Falls back gracefully when tokenizer has no EOS.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
@@ -237,6 +243,7 @@ def process_file(
|
||||
sentinel_ids=None,
|
||||
per_token=False,
|
||||
max_samples=None,
|
||||
eos_ids=None,
|
||||
):
|
||||
"""Score a single file, write per-sample JSONL, return summary stats."""
|
||||
if device is None:
|
||||
@@ -245,6 +252,11 @@ def process_file(
|
||||
if sentinel_ids is None:
|
||||
sentinel_ids = _resolve_sentinel_ids(tokenizer, "\n")
|
||||
|
||||
if eos_ids is None:
|
||||
eos_ids = []
|
||||
|
||||
eos_len = len(eos_ids)
|
||||
|
||||
data = _load_items(input_file)
|
||||
|
||||
if max_samples and len(data) > max_samples:
|
||||
@@ -267,7 +279,9 @@ def process_file(
|
||||
ctx_text = "\n\n".join(m["content"] for m in item["messages"][:i])
|
||||
ctx_ids = tokenizer.encode(ctx_text)
|
||||
resp_ids = tokenizer.encode(msg["content"], add_special_tokens=False)
|
||||
ctx_ids, resp_ids = _trim(ctx_ids, resp_ids, max_len)
|
||||
ctx_ids, resp_ids = _trim(ctx_ids, resp_ids, max_len - eos_len)
|
||||
if eos_ids and resp_ids and resp_ids[-1:] != eos_ids:
|
||||
resp_ids = resp_ids + eos_ids
|
||||
if ctx_ids and resp_ids:
|
||||
turns.append((ctx_ids, resp_ids))
|
||||
if not turns:
|
||||
@@ -284,7 +298,7 @@ def process_file(
|
||||
else:
|
||||
ctx_ids = tokenizer.encode(item[instr_key], add_special_tokens=False)
|
||||
resp_ids = tokenizer.encode(item[resp_key], add_special_tokens=False)
|
||||
ctx_ids, resp_ids = _trim(ctx_ids, resp_ids, max_len)
|
||||
ctx_ids, resp_ids = _trim(ctx_ids, resp_ids, max_len - eos_len)
|
||||
if not ctx_ids or not resp_ids:
|
||||
results.append(
|
||||
{
|
||||
@@ -294,6 +308,8 @@ def process_file(
|
||||
}
|
||||
)
|
||||
continue
|
||||
if eos_ids and resp_ids[-1:] != eos_ids:
|
||||
resp_ids = resp_ids + eos_ids
|
||||
buffer.append((item, [(ctx_ids, resp_ids)], "plain"))
|
||||
|
||||
if len(buffer) >= batch_size:
|
||||
@@ -452,6 +468,11 @@ def main():
|
||||
default=None,
|
||||
help="Maximum number of samples per file (random subsample). Default: all.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--append_eos/--no-append_eos",
|
||||
default=True,
|
||||
help="Append EOS token at the end of response in both passes (default: enabled).",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.device is None:
|
||||
@@ -466,6 +487,16 @@ def main():
|
||||
|
||||
sentinel_ids = _resolve_sentinel_ids(tokenizer, args.sentinel_text)
|
||||
|
||||
eos_ids = []
|
||||
if args.append_eos:
|
||||
eos_token_id = getattr(tokenizer, "eos_token_id", None)
|
||||
if eos_token_id is not None:
|
||||
eos_ids = [eos_token_id]
|
||||
else:
|
||||
print(
|
||||
"Warning: --append_eos enabled but tokenizer has no EOS token; skipping."
|
||||
)
|
||||
|
||||
input_files = _collect_input_files(args.input_path)
|
||||
if not input_files:
|
||||
print(f"No input files found at {args.input_path}")
|
||||
@@ -493,6 +524,7 @@ def main():
|
||||
sentinel_ids=sentinel_ids,
|
||||
per_token=args.per_token,
|
||||
max_samples=args.max_samples,
|
||||
eos_ids=eos_ids,
|
||||
)
|
||||
all_stats[label] = stats
|
||||
|
||||
|
||||
@@ -509,7 +509,10 @@ def main():
|
||||
help="Number of samples per problem (best-of-n scoring)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size", type=int, default=1, help="Inference batch size"
|
||||
"--batch_size", type=int, default=64, help="Inference batch size"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_seq_len", type=int, default=4096, help="Max sequence length for KV cache"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--limit",
|
||||
@@ -542,6 +545,7 @@ def main():
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=args.batch_size,
|
||||
max_seq_len=args.max_seq_len,
|
||||
)
|
||||
|
||||
results = evaluate(
|
||||
|
||||
@@ -179,31 +179,53 @@ def apply_chat(
|
||||
)
|
||||
|
||||
|
||||
def choice_logprob(
|
||||
model, tokenizer, context_ids: list[int], choice_letter: str, device: str
|
||||
) -> float:
|
||||
choice_text = choice_letter
|
||||
choice_ids = tokenizer.encode(choice_text, add_special_tokens=False)
|
||||
input_ids = context_ids + choice_ids
|
||||
max_len = model.config.max_position_embeddings
|
||||
if len(input_ids) > max_len:
|
||||
overflow = len(input_ids) - max_len
|
||||
input_ids = input_ids[overflow:]
|
||||
ctx_len = len(input_ids) - len(choice_ids)
|
||||
else:
|
||||
ctx_len = len(context_ids)
|
||||
def choice_logprobs_batched(
|
||||
model,
|
||||
tokenizer,
|
||||
context_ids_list: list[list[int]],
|
||||
device: str,
|
||||
max_model_len: int,
|
||||
) -> list[dict[str, float]]:
|
||||
"""Compute log-probs for multiple questions x 4 choices in batches.
|
||||
|
||||
Returns a list of dicts: [{A: score, B: score, C: score, D: score}, ...]
|
||||
"""
|
||||
letters = ("A", "B", "C", "D")
|
||||
choice_ids_list = [tokenizer.encode(c, add_special_tokens=False) for c in letters]
|
||||
|
||||
all_inputs: list[tuple[int, int, list[int], int, list[int]]] = []
|
||||
for qi, context_ids in enumerate(context_ids_list):
|
||||
for ci, choice_ids in enumerate(choice_ids_list):
|
||||
input_ids = context_ids + choice_ids
|
||||
if len(input_ids) > max_model_len:
|
||||
overflow = len(input_ids) - max_model_len
|
||||
input_ids = input_ids[overflow:]
|
||||
ctx_len = len(input_ids) - len(choice_ids)
|
||||
else:
|
||||
ctx_len = len(context_ids)
|
||||
all_inputs.append((qi, ci, input_ids, ctx_len, choice_ids))
|
||||
|
||||
n = len(all_inputs)
|
||||
max_input_len = max(len(x[2]) for x in all_inputs)
|
||||
padded = torch.zeros(n, max_input_len, dtype=torch.long, device=device)
|
||||
mask = torch.zeros(n, max_input_len, dtype=torch.bool, device=device)
|
||||
for i, (_, _, ids, _, _) in enumerate(all_inputs):
|
||||
padded[i, : len(ids)] = torch.tensor(ids, dtype=torch.long, device=device)
|
||||
mask[i, : len(ids)] = True
|
||||
|
||||
input_tensor = torch.tensor([input_ids], device=device, dtype=torch.long)
|
||||
with torch.inference_mode():
|
||||
logits = model(input_tensor)["logits"][0]
|
||||
logits = model(padded, input_mask=mask)["logits"]
|
||||
|
||||
score = 0.0
|
||||
for i, tid in enumerate(choice_ids):
|
||||
pos = ctx_len - 1 + i
|
||||
if pos >= len(logits):
|
||||
break
|
||||
score += F.log_softmax(logits[pos], dim=-1)[tid].item()
|
||||
return score
|
||||
results = [{} for _ in range(len(context_ids_list))]
|
||||
for i, (qi, ci, _, ctx_len, choice_ids) in enumerate(all_inputs):
|
||||
score = 0.0
|
||||
for j, tid in enumerate(choice_ids):
|
||||
pos = ctx_len - 1 + j
|
||||
if pos >= logits.size(1):
|
||||
break
|
||||
score += F.log_softmax(logits[i, pos].float(), dim=-1)[tid].item()
|
||||
results[qi][letters[ci]] = score
|
||||
return results
|
||||
|
||||
|
||||
def _permute_choices(item: dict, rng: random.Random) -> tuple[dict, str]:
|
||||
@@ -233,25 +255,42 @@ def evaluate_subject(
|
||||
device: str,
|
||||
n_shot: int,
|
||||
seed: int = 0,
|
||||
batch_size: int = 16,
|
||||
) -> tuple[float, int, int]:
|
||||
rng = random.Random(seed) if seed >= 0 else None
|
||||
correct = 0
|
||||
total = 0
|
||||
for item in tqdm.tqdm(test_data, desc=f"{subject:40s}", leave=False):
|
||||
|
||||
context_ids_list = []
|
||||
answers = []
|
||||
for item in test_data:
|
||||
if rng is not None:
|
||||
permuted, answer = _permute_choices(item, rng)
|
||||
else:
|
||||
permuted, answer = item, item["answer"]
|
||||
raw_prompt = build_prompt(permuted["question"], permuted, subject)
|
||||
context = apply_chat(tokenizer, raw_prompt, n_shot, dev_data or [], subject)
|
||||
context_ids = tokenizer.encode(context)
|
||||
scores = {
|
||||
c: choice_logprob(model, tokenizer, context_ids, c, device)
|
||||
for c in ("A", "B", "C", "D")
|
||||
}
|
||||
if max(scores, key=scores.get) == answer:
|
||||
correct += 1
|
||||
total += 1
|
||||
context_ids_list.append(tokenizer.encode(context))
|
||||
answers.append(answer)
|
||||
|
||||
max_model_len = model.config.max_position_embeddings
|
||||
|
||||
num_batches = (len(context_ids_list) + batch_size - 1) // batch_size
|
||||
for start in tqdm.tqdm(
|
||||
range(0, len(context_ids_list), batch_size),
|
||||
total=num_batches,
|
||||
desc=f"{subject:40s}",
|
||||
leave=False,
|
||||
):
|
||||
batch = context_ids_list[start : start + batch_size]
|
||||
batch_answers = answers[start : start + batch_size]
|
||||
scores_list = choice_logprobs_batched(
|
||||
model, tokenizer, batch, device, max_model_len
|
||||
)
|
||||
for scores, answer in zip(scores_list, batch_answers):
|
||||
if max(scores, key=scores.get) == answer:
|
||||
correct += 1
|
||||
total += 1
|
||||
return correct / total, correct, total
|
||||
|
||||
|
||||
@@ -290,6 +329,12 @@ def main():
|
||||
default=0,
|
||||
help="Seed for option permutation (0 to enable, -1 to disable)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size",
|
||||
type=int,
|
||||
default=4,
|
||||
help="Number of questions per batch (4 choices each = 4*B rows)",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.download or not os.path.exists(args.data_dir):
|
||||
@@ -329,6 +374,7 @@ def main():
|
||||
device,
|
||||
args.n_shot,
|
||||
seed=args.seed,
|
||||
batch_size=args.batch_size,
|
||||
)
|
||||
results[subject] = {"accuracy": round(acc, 4), "correct": corr, "total": tot}
|
||||
total_correct += corr
|
||||
|
||||
@@ -415,7 +415,7 @@ if __name__ == "__main__":
|
||||
help="Key for the text field in the input data.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size", type=int, default=4, help="Batch size for evaluation."
|
||||
"--batch_size", type=int, default=64, help="Batch size for evaluation."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_length",
|
||||
|
||||
+214
-236
@@ -1,299 +1,277 @@
|
||||
"""Benchmark AutoRegressiveLM with KVCache"""
|
||||
|
||||
import argparse
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import click
|
||||
import torch
|
||||
|
||||
from astrai import setup_logging
|
||||
from astrai.config import AutoRegressiveLMConfig
|
||||
from astrai.inference import ContiguousCache, PageCache
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
from astrai.extension import ATTN_BACKEND, attn_backend
|
||||
from astrai.inference.core.cache import PagePool
|
||||
from astrai.model import AutoModel
|
||||
|
||||
_DTYPES = ["bfloat16", "float16", "float32"]
|
||||
_CACHES = ["contiguous", "paged"]
|
||||
DEFAULT_CKPT = str(Path(__file__).resolve().parents[2] / "ckpt_bucket" / "kami-15bt")
|
||||
CACHE_MAX_SEQ = 2048
|
||||
|
||||
|
||||
@dataclass
|
||||
class BenchmarkResult:
|
||||
total_tokens: int
|
||||
total_time: float
|
||||
tokens_per_second: float
|
||||
metadata: Dict[str, Any]
|
||||
def __init__(
|
||||
self,
|
||||
name: str,
|
||||
batch_size: int,
|
||||
seq_len: int,
|
||||
tokens_per_second: float,
|
||||
latency_ms: float,
|
||||
metadata: Optional[dict] = None,
|
||||
):
|
||||
self.name = name
|
||||
self.batch_size = batch_size
|
||||
self.seq_len = seq_len
|
||||
self.tokens_per_second = tokens_per_second
|
||||
self.latency_ms = latency_ms
|
||||
self.metadata = metadata or {}
|
||||
|
||||
|
||||
class GenerationBenchmark:
|
||||
def __init__(
|
||||
self,
|
||||
model: AutoModel,
|
||||
config: AutoRegressiveLMConfig,
|
||||
device: str = "cuda",
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
cache_type: str = "contiguous",
|
||||
):
|
||||
self.config = config
|
||||
self.device = device
|
||||
self.dtype = dtype
|
||||
self.cache_type = cache_type
|
||||
self.model = AutoRegressiveLM(config).to(device=device, dtype=dtype)
|
||||
self.model.eval()
|
||||
self.model = model
|
||||
self.config = config
|
||||
|
||||
@torch.inference_mode()
|
||||
def _make_pool(self, batch_size: int) -> PagePool:
|
||||
return PagePool(
|
||||
n_layers=self.config.num_hidden_layers,
|
||||
n_kv_heads=self.config.num_key_value_heads,
|
||||
head_dim=self.config.hidden_size // self.config.num_attention_heads,
|
||||
max_batch_size=batch_size,
|
||||
max_seq_len=CACHE_MAX_SEQ,
|
||||
device=self.device,
|
||||
dtype=self.dtype,
|
||||
page_size=1,
|
||||
n_tokens=None,
|
||||
)
|
||||
|
||||
def _run_prefill(self, pool: PagePool, batch_size: int, prompt_len: int) -> list:
|
||||
input_ids = torch.randint(
|
||||
0, self.config.vocab_size, (batch_size, prompt_len), device=self.device
|
||||
)
|
||||
position_ids = (
|
||||
torch.arange(0, prompt_len, dtype=torch.long, device=self.device)
|
||||
.unsqueeze(0)
|
||||
.expand(batch_size, -1)
|
||||
)
|
||||
input_mask = position_ids.unsqueeze(-1) >= torch.arange(
|
||||
prompt_len, device=self.device
|
||||
)
|
||||
|
||||
task_ids = [f"bench_{i}" for i in range(batch_size)]
|
||||
for tid in task_ids:
|
||||
pool.task_alloc(tid, list(range(prompt_len)))
|
||||
|
||||
kv_cache = pool.bind_tasks(
|
||||
task_ids, [prompt_len] * batch_size, self.device, start_pos=0
|
||||
)
|
||||
with torch.inference_mode(), attn_backend(ATTN_BACKEND.CUDA):
|
||||
self.model(
|
||||
input_ids,
|
||||
input_mask=input_mask,
|
||||
kv_cache=kv_cache,
|
||||
position_ids=position_ids,
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
return task_ids
|
||||
|
||||
def _run_decode_step(self, pool: PagePool, task_ids: list, seq_len: int):
|
||||
batch_size = len(task_ids)
|
||||
input_ids = torch.randint(
|
||||
0, self.config.vocab_size, (batch_size, 1), device=self.device
|
||||
)
|
||||
position_ids = torch.tensor(
|
||||
[[seq_len] for _ in range(batch_size)], dtype=torch.long, device=self.device
|
||||
)
|
||||
total_len = seq_len + 1
|
||||
input_mask = position_ids[:, :, None] >= torch.arange(
|
||||
total_len, device=self.device
|
||||
)
|
||||
kv_cache = pool.bind_tasks(task_ids, [seq_len + 1] * batch_size, self.device)
|
||||
with torch.inference_mode(), attn_backend(ATTN_BACKEND.CUDA):
|
||||
self.model(
|
||||
input_ids,
|
||||
input_mask=input_mask,
|
||||
kv_cache=kv_cache,
|
||||
position_ids=position_ids,
|
||||
)
|
||||
|
||||
def run_prefill_benchmark(
|
||||
self,
|
||||
batch_size: int = 1,
|
||||
batch_size: int = 4,
|
||||
prompt_length: int = 512,
|
||||
num_trials: int = 10,
|
||||
num_trials: int = 5,
|
||||
) -> BenchmarkResult:
|
||||
for _ in range(3):
|
||||
prompt_ids = torch.randint(
|
||||
0,
|
||||
self.config.vocab_size,
|
||||
(batch_size, prompt_length),
|
||||
device=self.device,
|
||||
dtype=torch.long,
|
||||
)
|
||||
_ = self.model(prompt_ids)
|
||||
torch.cuda.synchronize()
|
||||
import time
|
||||
|
||||
total_time = 0.0
|
||||
total_tokens = batch_size * prompt_length * num_trials
|
||||
input_ids = torch.randint(
|
||||
0, self.config.vocab_size, (batch_size, prompt_length), device=self.device
|
||||
)
|
||||
position_ids = (
|
||||
torch.arange(0, prompt_length, dtype=torch.long, device=self.device)
|
||||
.unsqueeze(0)
|
||||
.expand(batch_size, -1)
|
||||
)
|
||||
|
||||
for trial in range(num_trials):
|
||||
prompt_ids = torch.randint(
|
||||
0,
|
||||
self.config.vocab_size,
|
||||
(batch_size, prompt_length),
|
||||
device=self.device,
|
||||
dtype=torch.long,
|
||||
)
|
||||
start = torch.cuda.Event(enable_timing=True)
|
||||
end = torch.cuda.Event(enable_timing=True)
|
||||
start.record()
|
||||
_ = self.model(prompt_ids)
|
||||
end.record()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
trial_time = start.elapsed_time(end) / 1000
|
||||
total_time += trial_time
|
||||
|
||||
print(
|
||||
f" Trial {trial + 1}/{num_trials}: {prompt_length} tokens in {trial_time:.3f}s "
|
||||
f"({prompt_length / trial_time:.1f} tok/s)"
|
||||
)
|
||||
for _ in range(3):
|
||||
with torch.inference_mode(), attn_backend(ATTN_BACKEND.CUDA):
|
||||
self.model(input_ids, position_ids=position_ids)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
t0 = time.perf_counter()
|
||||
for _ in range(num_trials):
|
||||
with torch.inference_mode(), attn_backend(ATTN_BACKEND.CUDA):
|
||||
self.model(input_ids, position_ids=position_ids)
|
||||
torch.cuda.synchronize()
|
||||
elapsed = time.perf_counter() - t0
|
||||
tokens = batch_size * prompt_length * num_trials
|
||||
tps = tokens / elapsed
|
||||
return BenchmarkResult(
|
||||
total_tokens=total_tokens,
|
||||
total_time=total_time,
|
||||
tokens_per_second=total_tokens / total_time,
|
||||
metadata={
|
||||
"benchmark_type": "prefill",
|
||||
"batch_size": batch_size,
|
||||
"prompt_length": prompt_length,
|
||||
"dtype": str(self.dtype),
|
||||
"device": self.device,
|
||||
"cache": "none",
|
||||
},
|
||||
name="prefill",
|
||||
batch_size=batch_size,
|
||||
seq_len=prompt_length,
|
||||
tokens_per_second=tps,
|
||||
latency_ms=elapsed / num_trials * 1000,
|
||||
metadata={"benchmark_type": "prefill", "num_trials": num_trials},
|
||||
)
|
||||
|
||||
@torch.inference_mode()
|
||||
def run_decoding_benchmark(
|
||||
self,
|
||||
batch_size: int = 1,
|
||||
batch_size: int = 4,
|
||||
prompt_length: int = 512,
|
||||
gen_length: int = 128,
|
||||
num_trials: int = 5,
|
||||
) -> BenchmarkResult:
|
||||
total_time = 0.0
|
||||
total_tokens = batch_size * gen_length * num_trials
|
||||
import time
|
||||
|
||||
for trial in range(num_trials):
|
||||
prompt_ids = torch.randint(
|
||||
0,
|
||||
self.config.vocab_size,
|
||||
(batch_size, prompt_length),
|
||||
device=self.device,
|
||||
dtype=torch.long,
|
||||
)
|
||||
gen_ids = torch.randint(
|
||||
0,
|
||||
self.config.vocab_size,
|
||||
(batch_size, gen_length),
|
||||
device=self.device,
|
||||
dtype=torch.long,
|
||||
)
|
||||
pool = self._make_pool(batch_size)
|
||||
task_ids = self._run_prefill(pool, batch_size, prompt_length)
|
||||
|
||||
head_dim = self.config.hidden_size // self.config.num_attention_heads
|
||||
max_seq = prompt_length + gen_length
|
||||
|
||||
if self.cache_type == "contiguous":
|
||||
cache = ContiguousCache(
|
||||
self.config.num_hidden_layers,
|
||||
batch_size,
|
||||
max_seq,
|
||||
self.config.num_key_value_heads,
|
||||
head_dim,
|
||||
self.device,
|
||||
self.dtype,
|
||||
)
|
||||
else:
|
||||
page_size = 128
|
||||
n_pages = (max_seq + page_size - 1) // page_size * batch_size
|
||||
cache = PageCache(
|
||||
self.config.num_hidden_layers,
|
||||
n_pages,
|
||||
page_size,
|
||||
self.config.num_key_value_heads,
|
||||
head_dim,
|
||||
self.device,
|
||||
self.dtype,
|
||||
)
|
||||
|
||||
task_ids = [f"b{i}" for i in range(batch_size)]
|
||||
for tid in task_ids:
|
||||
cache.task_alloc(tid, [0] * max_seq)
|
||||
for p in range(max_seq):
|
||||
cache.task_extend(tid, p)
|
||||
|
||||
cv = cache.bind_tasks(task_ids, prompt_length, self.device)
|
||||
_ = self.model(
|
||||
prompt_ids,
|
||||
paged_cache=cv,
|
||||
position_ids=torch.arange(
|
||||
prompt_length, dtype=torch.long, device=self.device
|
||||
)
|
||||
.unsqueeze(0)
|
||||
.expand(batch_size, -1),
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
start = torch.cuda.Event(enable_timing=True)
|
||||
end = torch.cuda.Event(enable_timing=True)
|
||||
start.record()
|
||||
|
||||
for i in range(gen_length):
|
||||
pos = prompt_length + i
|
||||
cv = cache.bind_tasks(task_ids, pos + 1, self.device)
|
||||
_ = self.model(
|
||||
gen_ids[:, i : i + 1],
|
||||
paged_cache=cv,
|
||||
position_ids=torch.full(
|
||||
(batch_size, 1),
|
||||
pos,
|
||||
dtype=torch.long,
|
||||
device=self.device,
|
||||
),
|
||||
)
|
||||
|
||||
end.record()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
for tid in task_ids:
|
||||
cache.task_free(tid)
|
||||
|
||||
trial_time = start.elapsed_time(end) / 1000
|
||||
total_time += trial_time
|
||||
|
||||
print(
|
||||
f" Trial {trial + 1}/{num_trials}: {gen_length} tokens in {trial_time:.3f}s "
|
||||
f"({gen_length / trial_time:.1f} tok/s)"
|
||||
)
|
||||
for i in range(5):
|
||||
self._run_decode_step(pool, task_ids, prompt_length + i)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
t0 = time.perf_counter()
|
||||
for i in range(gen_length * num_trials):
|
||||
self._run_decode_step(pool, task_ids, prompt_length + 5 + i)
|
||||
torch.cuda.synchronize()
|
||||
elapsed = time.perf_counter() - t0
|
||||
tokens = batch_size * gen_length * num_trials
|
||||
tps = tokens / elapsed
|
||||
return BenchmarkResult(
|
||||
total_tokens=total_tokens,
|
||||
total_time=total_time,
|
||||
tokens_per_second=total_tokens / total_time,
|
||||
name="decode",
|
||||
batch_size=batch_size,
|
||||
seq_len=gen_length,
|
||||
tokens_per_second=tps,
|
||||
latency_ms=elapsed / (gen_length * num_trials) * 1000,
|
||||
metadata={
|
||||
"benchmark_type": "decoding",
|
||||
"batch_size": batch_size,
|
||||
"benchmark_type": "decode",
|
||||
"num_trials": num_trials,
|
||||
"prompt_length": prompt_length,
|
||||
"gen_length": gen_length,
|
||||
"dtype": str(self.dtype),
|
||||
"device": self.device,
|
||||
"cache": self.cache_type,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def print_benchmark_result(result: BenchmarkResult):
|
||||
btype = result.metadata["benchmark_type"]
|
||||
print(f"\n{' ' + btype.upper() + ' Benchmark ':-^80}")
|
||||
print(f"Total Tokens Processed: {result.total_tokens:,}")
|
||||
print(f"Time Consumed: {result.total_time:.3f}s")
|
||||
print(f"Throughput: {result.tokens_per_second:,.1f} tok/s")
|
||||
def print_benchmark_result(result: BenchmarkResult) -> None:
|
||||
print("-" * 80)
|
||||
print(f"{result.name.upper()} — Batch={result.batch_size}, SeqLen={result.seq_len}")
|
||||
print(f" Throughput : {result.tokens_per_second:.1f} tokens/s")
|
||||
print(f" Latency : {result.latency_ms:.2f} ms/step")
|
||||
for k, v in result.metadata.items():
|
||||
if k != "benchmark_type":
|
||||
print(f"{k.replace('_', ' ').title()}: {v}")
|
||||
print(f" {k.replace('_', ' ').title()}: {v}")
|
||||
print("-" * 80)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="AutoRegressiveLM benchmark")
|
||||
parser.add_argument(
|
||||
"--device", type=str, default="cuda", help="Device (default: cuda)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dtype",
|
||||
type=str,
|
||||
default="bfloat16",
|
||||
choices=["bfloat16", "float16", "float32"],
|
||||
help="Dtype",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--cache",
|
||||
type=str,
|
||||
default="contiguous",
|
||||
choices=["contiguous", "paged"],
|
||||
help="KV cache type",
|
||||
)
|
||||
parser.add_argument("--batch_size", type=int, default=4, help="Batch size")
|
||||
parser.add_argument("--prompt_length", type=int, default=512, help="Prompt length")
|
||||
parser.add_argument("--gen_length", type=int, default=128, help="Generation length")
|
||||
parser.add_argument("--num_trials", type=int, default=5, help="Number of trials")
|
||||
parser.add_argument(
|
||||
"--prefill_only", action="store_true", help="Run prefill benchmark only"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--decode_only", action="store_true", help="Run decoding benchmark only"
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
dtype_map = {
|
||||
@click.command(name="benchmark", help="Benchmark model throughput and latency.")
|
||||
@click.option("--device", default="cuda", help="Device.")
|
||||
@click.option(
|
||||
"--dtype", type=click.Choice(_DTYPES), default="bfloat16", help="Data type."
|
||||
)
|
||||
@click.option(
|
||||
"--cache", type=click.Choice(_CACHES), default="contiguous", help="KV cache type."
|
||||
)
|
||||
@click.option("--batch_size", type=int, default=4, help="Batch size.")
|
||||
@click.option("--prompt_length", type=int, default=512, help="Prompt length.")
|
||||
@click.option("--gen_length", type=int, default=128, help="Generation length.")
|
||||
@click.option("--num_trials", type=int, default=5, help="Number of trials.")
|
||||
@click.option("--prefill_only", is_flag=True, help="Prefill benchmark only.")
|
||||
@click.option("--decode_only", is_flag=True, help="Decode benchmark only.")
|
||||
@click.option(
|
||||
"--ckpt",
|
||||
default=DEFAULT_CKPT,
|
||||
help="Checkpoint directory.",
|
||||
)
|
||||
def benchmark_command(
|
||||
device: str,
|
||||
dtype: str,
|
||||
cache: str,
|
||||
batch_size: int,
|
||||
prompt_length: int,
|
||||
gen_length: int,
|
||||
num_trials: int,
|
||||
prefill_only: bool,
|
||||
decode_only: bool,
|
||||
ckpt: str,
|
||||
) -> None:
|
||||
"""Benchmark model throughput and latency."""
|
||||
dtype_map: dict[str, torch.dtype] = {
|
||||
"bfloat16": torch.bfloat16,
|
||||
"float16": torch.float16,
|
||||
"float32": torch.float32,
|
||||
}
|
||||
|
||||
config = AutoRegressiveLMConfig(
|
||||
vocab_size=10000,
|
||||
hidden_size=1536,
|
||||
num_attention_heads=24,
|
||||
num_key_value_heads=4,
|
||||
intermediate_size=6912,
|
||||
max_position_embeddings=2048,
|
||||
num_hidden_layers=24,
|
||||
rms_norm_eps=1e-5,
|
||||
click.echo(f"Loading model from {ckpt} ...")
|
||||
config = AutoRegressiveLMConfig.from_file(str(Path(ckpt) / "config.json"))
|
||||
model = AutoModel.from_pretrained(ckpt)
|
||||
model.to(device=device, dtype=dtype_map[dtype])
|
||||
model.eval()
|
||||
|
||||
bench = GenerationBenchmark(
|
||||
model=model,
|
||||
config=config,
|
||||
device=device,
|
||||
dtype=dtype_map[dtype],
|
||||
cache_type=cache,
|
||||
)
|
||||
|
||||
benchmark = GenerationBenchmark(
|
||||
config, device=args.device, dtype=dtype_map[args.dtype], cache_type=args.cache
|
||||
)
|
||||
click.secho(f"Benchmark: device={device} dtype={dtype}", bold=True)
|
||||
|
||||
print("=" * 80)
|
||||
print(
|
||||
f"Running AutoRegressiveLM Benchmark (device={args.device}, dtype={args.dtype})"
|
||||
)
|
||||
print("=" * 80)
|
||||
|
||||
if not args.decode_only:
|
||||
prefill_result = benchmark.run_prefill_benchmark(
|
||||
batch_size=args.batch_size,
|
||||
prompt_length=args.prompt_length,
|
||||
num_trials=args.num_trials,
|
||||
if not decode_only:
|
||||
result = bench.run_prefill_benchmark(
|
||||
batch_size=batch_size,
|
||||
prompt_length=prompt_length,
|
||||
num_trials=num_trials,
|
||||
)
|
||||
print_benchmark_result(prefill_result)
|
||||
print_benchmark_result(result)
|
||||
|
||||
if not args.prefill_only:
|
||||
gen_result = benchmark.run_decoding_benchmark(
|
||||
batch_size=args.batch_size,
|
||||
prompt_length=args.prompt_length,
|
||||
gen_length=args.gen_length,
|
||||
num_trials=args.num_trials,
|
||||
if not prefill_only:
|
||||
result = bench.run_decoding_benchmark(
|
||||
batch_size=batch_size,
|
||||
prompt_length=prompt_length,
|
||||
gen_length=gen_length,
|
||||
num_trials=num_trials,
|
||||
)
|
||||
print_benchmark_result(gen_result)
|
||||
print_benchmark_result(result)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
setup_logging()
|
||||
benchmark_command()
|
||||
|
||||
+47
-101
@@ -1,11 +1,12 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
from typing import Optional
|
||||
|
||||
import click
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from astrai import setup_logging
|
||||
from astrai.inference import InferenceEngine
|
||||
from astrai.model import AutoModel
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
@@ -20,10 +21,9 @@ def processor(
|
||||
top_p: float,
|
||||
question_key: str,
|
||||
response_key: str,
|
||||
max_tokens: Optional[int],
|
||||
batch_size: int,
|
||||
num_samples: int = 1,
|
||||
cache_len: int = 2048,
|
||||
max_seq_len: Optional[int] = None,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
):
|
||||
@@ -38,8 +38,7 @@ def processor(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=batch_size * num_samples,
|
||||
max_seq_len=cache_len,
|
||||
max_prompt_len=cache_len,
|
||||
max_seq_len=max_seq_len,
|
||||
)
|
||||
|
||||
print(f"Reading {input_json_file} ...")
|
||||
@@ -55,9 +54,6 @@ def processor(
|
||||
prompts = [item[question_key] for item in input_data]
|
||||
print(f" {len(prompts)} prompts loaded\n")
|
||||
|
||||
if max_tokens is None:
|
||||
max_tokens = model.config.max_position_embeddings
|
||||
|
||||
chunk_size = max(1, batch_size)
|
||||
|
||||
with open(output_json_file, "w", encoding="utf-8") as f:
|
||||
@@ -74,7 +70,6 @@ def processor(
|
||||
resp_chunk = engine.generate(
|
||||
prompt=chunk_expanded,
|
||||
stream=False,
|
||||
max_tokens=max_tokens,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
@@ -89,7 +84,6 @@ def processor(
|
||||
resp_chunk = engine.generate(
|
||||
prompt=chunk,
|
||||
stream=False,
|
||||
max_tokens=max_tokens,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
@@ -121,95 +115,47 @@ def processor(
|
||||
engine.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="Batch generation from JSONL file.")
|
||||
|
||||
parser.add_argument(
|
||||
"--param_path", type=str, required=True, help="Path to the model directory."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--input_json_file",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to the input JSONL file.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output_json_file",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to the output JSONL file.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--question_key",
|
||||
type=str,
|
||||
default="question",
|
||||
help="Key for the question in the input JSON (default: question).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--response_key",
|
||||
type=str,
|
||||
default="response",
|
||||
help="Key for the response in the output JSON (default: response).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--temperature",
|
||||
type=float,
|
||||
default=0.60,
|
||||
help="Temperature for generating responses (default: 0.60).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--top_k",
|
||||
type=int,
|
||||
default=30,
|
||||
help="Top-k value for generating responses (default: 30).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--top_p",
|
||||
type=float,
|
||||
default=0.95,
|
||||
help="Top-p value for generating responses (default: 0.95).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Batch size for generating responses (default: 1).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num_samples",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of responses per prompt (expands batch internally, default: 1).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_tokens",
|
||||
type=int,
|
||||
default=None,
|
||||
help=(
|
||||
"Maximum tokens to generate "
|
||||
"(default: model config max_position_embeddings)."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--cache_len",
|
||||
type=int,
|
||||
default=2048,
|
||||
help="KV cache & prompt truncation length (default: 2048, lower = less memory).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--frequency_penalty",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="Frequency penalty to reduce repetition (default: 0.0, try 0.5-1.0).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--rep_window",
|
||||
type=int,
|
||||
default=64,
|
||||
help="Window size for frequency penalty (default: 64).",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
@click.command(name="generate", help="Batch generation from a JSONL prompt file.")
|
||||
@click.option(
|
||||
"--param_path",
|
||||
type=click.Path(exists=True),
|
||||
required=True,
|
||||
help="Path to the model directory.",
|
||||
)
|
||||
@click.option(
|
||||
"--input_json_file",
|
||||
type=click.Path(exists=True),
|
||||
required=True,
|
||||
help="Path to the input JSONL file.",
|
||||
)
|
||||
@click.option(
|
||||
"--output_json_file",
|
||||
type=click.Path(),
|
||||
required=True,
|
||||
help="Path to the output JSONL file.",
|
||||
)
|
||||
@click.option(
|
||||
"--question_key", default="question", help="Key for the question in input JSON."
|
||||
)
|
||||
@click.option(
|
||||
"--response_key", default="response", help="Key for the response in output JSON."
|
||||
)
|
||||
@click.option("--temperature", type=float, default=0.8, help="Sampling temperature.")
|
||||
@click.option("--top_k", type=int, default=50, help="Top-k filtering.")
|
||||
@click.option("--top_p", type=float, default=0.95, help="Top-p filtering.")
|
||||
@click.option("--batch_size", type=int, default=1, help="Batch size.")
|
||||
@click.option("--num_samples", type=int, default=1, help="Responses per prompt.")
|
||||
@click.option("--max_seq_len", type=int, default=2048, help="KV cache length.")
|
||||
@click.option("--frequency_penalty", type=float, default=0.0, help="Frequency penalty.")
|
||||
@click.option(
|
||||
"--rep_window", type=int, default=64, help="Window size for frequency penalty."
|
||||
)
|
||||
def generate_command(**kwargs):
|
||||
"""Batch generation from a JSONL prompt file."""
|
||||
with torch.inference_mode():
|
||||
processor(**vars(args))
|
||||
processor(**kwargs)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
setup_logging()
|
||||
generate_command()
|
||||
|
||||
+39
-35
@@ -1,48 +1,52 @@
|
||||
"""CLI: JSONL → tokenized .h5/.bin via config-driven Pipeline."""
|
||||
"""CLI: JSONL → tokenized .bin via config-driven Pipeline."""
|
||||
|
||||
import argparse
|
||||
import click
|
||||
|
||||
from astrai import setup_logging
|
||||
from astrai.config.preprocess_config import PipelineConfig
|
||||
from astrai.preprocessing.pipeline import Pipeline
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Raw JSONL → tokenized .h5/.bin via config-driven Pipeline"
|
||||
)
|
||||
parser.add_argument(
|
||||
"inputs", nargs="+", metavar="JSONL", help="One or more JSONL files"
|
||||
)
|
||||
parser.add_argument("--output_dir", "-o", required=True, help="Output directory")
|
||||
parser.add_argument(
|
||||
"--config", "-c", required=True, help="Path to pipeline config JSON"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--tokenizer_path",
|
||||
default="params",
|
||||
help="Path to tokenizer directory (default: params)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Number of records tokenized together (default: config value)",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
config = PipelineConfig.from_file(args.config)
|
||||
if args.batch_size is not None:
|
||||
if args.batch_size < 1:
|
||||
parser.error("--batch_size must be at least 1")
|
||||
config.preprocessing.batch_size = args.batch_size
|
||||
@click.command(
|
||||
name="preprocess", help="Tokenize and pack raw JSONL data into .bin format."
|
||||
)
|
||||
@click.argument("inputs", nargs=-1, type=click.Path(exists=True), required=True)
|
||||
@click.option(
|
||||
"--output_dir", "-o", type=click.Path(), required=True, help="Output directory."
|
||||
)
|
||||
@click.option(
|
||||
"--config",
|
||||
"-c",
|
||||
"pipeline_config",
|
||||
type=click.Path(exists=True),
|
||||
required=True,
|
||||
help="Pipeline config JSON.",
|
||||
)
|
||||
@click.option(
|
||||
"--tokenizer_path",
|
||||
type=click.Path(exists=True),
|
||||
default="params",
|
||||
help="Path to tokenizer directory.",
|
||||
)
|
||||
@click.option("--batch_size", type=int, default=None, help="Records per batch.")
|
||||
def preprocess_command(inputs, output_dir, pipeline_config, tokenizer_path, batch_size):
|
||||
"""Tokenize and pack raw JSONL data into .bin format."""
|
||||
config = PipelineConfig.from_file(pipeline_config)
|
||||
if batch_size is not None:
|
||||
if batch_size < 1:
|
||||
raise click.BadParameter("--batch_size must be at least 1")
|
||||
config.preprocessing.batch_size = batch_size
|
||||
|
||||
click.echo(f"Preprocessing {len(inputs)} file(s) → {output_dir}")
|
||||
Pipeline(
|
||||
config=config,
|
||||
input_paths=args.inputs,
|
||||
output_dir=args.output_dir,
|
||||
tokenizer_path=args.tokenizer_path,
|
||||
input_paths=list(inputs),
|
||||
output_dir=output_dir,
|
||||
tokenizer_path=tokenizer_path,
|
||||
).run()
|
||||
click.echo("Done.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
setup_logging()
|
||||
preprocess_command()
|
||||
|
||||
+50
-53
@@ -1,72 +1,69 @@
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
import click
|
||||
import torch
|
||||
|
||||
from astrai import setup_logging
|
||||
from astrai.inference import run_server
|
||||
|
||||
_DTYPES = ["bfloat16", "float16", "float32"]
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Start AstrAI inference HTTP server")
|
||||
parser.add_argument(
|
||||
"--host", default="0.0.0.0", help="Host address (default: 0.0.0.0)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--port", type=int, default=8000, help="Port number (default: 8000)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--reload", action="store_true", help="Enable auto-reload for development"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--param_path",
|
||||
type=Path,
|
||||
default=None,
|
||||
help="Path to model parameters (default: project_root/params)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--device",
|
||||
type=str,
|
||||
default="cuda",
|
||||
help="Device to load model on (default: cuda)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dtype",
|
||||
type=str,
|
||||
default="bfloat16",
|
||||
choices=["bfloat16", "float16", "float32"],
|
||||
help="Data type for model weights (default: bfloat16)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_batch_size",
|
||||
type=int,
|
||||
default=16,
|
||||
help="Maximum batch size for continuous batching (default: 16)",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
# Convert dtype string to torch dtype
|
||||
@click.command(name="serve", help="Launch inference server (OpenAI-compatible API).")
|
||||
@click.option("--host", default="0.0.0.0", help="Host address.")
|
||||
@click.option("--port", type=int, default=8000, help="Port number.")
|
||||
@click.option("--reload", is_flag=True, help="Enable auto-reload for development.")
|
||||
@click.option(
|
||||
"--param_path",
|
||||
type=click.Path(exists=True),
|
||||
default=None,
|
||||
help="Path to model parameters.",
|
||||
)
|
||||
@click.option("--device", default="cuda", help="Device to load model on.")
|
||||
@click.option(
|
||||
"--dtype",
|
||||
type=click.Choice(_DTYPES),
|
||||
default="bfloat16",
|
||||
help="Data type for model weights.",
|
||||
)
|
||||
@click.option(
|
||||
"--max_batch_size",
|
||||
type=int,
|
||||
default=16,
|
||||
help="Maximum batch size for continuous batching.",
|
||||
)
|
||||
@click.option(
|
||||
"--max_seq_len",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Maximum sequence length (KV cache size + prompt truncation). Uses model config if not set.",
|
||||
)
|
||||
def server_command(
|
||||
host, port, reload, param_path, device, dtype, max_batch_size, max_seq_len
|
||||
):
|
||||
"""Launch inference server (OpenAI-compatible API)."""
|
||||
dtype_map = {
|
||||
"bfloat16": torch.bfloat16,
|
||||
"float16": torch.float16,
|
||||
"float32": torch.float32,
|
||||
}
|
||||
dtype = dtype_map[args.dtype]
|
||||
|
||||
project_root = Path(__file__).parent.parent.parent
|
||||
param_path = args.param_path or (project_root / "params")
|
||||
print(f"Starting AstrAI inference server on http://{args.host}:{args.port}")
|
||||
print(f"Model parameters expected at: {param_path}")
|
||||
print(f"Device: {args.device}, Dtype: {args.dtype}")
|
||||
param_path = param_path or str(project_root / "params")
|
||||
|
||||
click.echo(f"Starting server on http://{host}:{port}")
|
||||
click.echo(f"Model: {param_path} | Device: {device} | Dtype: {dtype}")
|
||||
run_server(
|
||||
host=args.host,
|
||||
port=args.port,
|
||||
reload=args.reload,
|
||||
device=args.device,
|
||||
dtype=dtype,
|
||||
param_path=param_path,
|
||||
max_batch_size=args.max_batch_size,
|
||||
host=host,
|
||||
port=port,
|
||||
reload=reload,
|
||||
device=device,
|
||||
dtype=dtype_map[dtype],
|
||||
param_path=Path(param_path),
|
||||
max_batch_size=max_batch_size,
|
||||
max_seq_len=max_seq_len,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
setup_logging()
|
||||
server_command()
|
||||
|
||||
+253
-308
@@ -1,12 +1,13 @@
|
||||
import argparse
|
||||
import os
|
||||
from collections.abc import Callable
|
||||
from functools import partial
|
||||
from typing import Any, Callable, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
import click
|
||||
import torch
|
||||
import torch.optim as optim
|
||||
from torch import Tensor, nn
|
||||
from torch import Tensor, nn, optim
|
||||
|
||||
from astrai import setup_logging
|
||||
from astrai.config import AutoRegressiveLMConfig, TrainConfig
|
||||
from astrai.dataset import DatasetFactory, dpo_collate_fn, grpo_collate_fn
|
||||
from astrai.model import AutoRegressiveLM
|
||||
@@ -28,14 +29,14 @@ class MuonMix(optim.Optimizer):
|
||||
ns_steps: int = 5,
|
||||
adjust_lr_fn: str = "match_rms_adamw",
|
||||
):
|
||||
defaults = dict(
|
||||
lr=lr,
|
||||
weight_decay=weight_decay,
|
||||
momentum=momentum,
|
||||
nesterov=nesterov,
|
||||
ns_steps=ns_steps,
|
||||
adjust_lr_fn=adjust_lr_fn,
|
||||
)
|
||||
defaults = {
|
||||
"lr": lr,
|
||||
"weight_decay": weight_decay,
|
||||
"momentum": momentum,
|
||||
"nesterov": nesterov,
|
||||
"ns_steps": ns_steps,
|
||||
"adjust_lr_fn": adjust_lr_fn,
|
||||
}
|
||||
params = [p for p in model.parameters() if p.requires_grad]
|
||||
super().__init__(params, defaults)
|
||||
|
||||
@@ -82,312 +83,253 @@ class MuonMix(optim.Optimizer):
|
||||
self.muon.zero_grad(set_to_none)
|
||||
self.adamw.zero_grad(set_to_none)
|
||||
|
||||
def state_dict(self) -> Dict[str, Any]:
|
||||
def state_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"muon": self.muon.state_dict(),
|
||||
"adamw": self.adamw.state_dict(),
|
||||
}
|
||||
|
||||
def load_state_dict(self, state_dict: Dict[str, Any]):
|
||||
def load_state_dict(self, state_dict: dict[str, Any]):
|
||||
self.muon.load_state_dict(state_dict["muon"])
|
||||
self.adamw.load_state_dict(state_dict["adamw"])
|
||||
self.param_groups = [*self.muon.param_groups, *self.adamw.param_groups]
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
def _merge_yaml_into_kwargs(config_path: str, passed_kwargs: dict) -> dict:
|
||||
"""Load YAML config, then override with explicit CLI kwargs (None excluded)."""
|
||||
import yaml
|
||||
|
||||
parser = argparse.ArgumentParser(description="Train the AutoRegressiveLM model.")
|
||||
with open(config_path) as f:
|
||||
cfg = yaml.safe_load(f)
|
||||
|
||||
parser.add_argument(
|
||||
"--train_type",
|
||||
type=str,
|
||||
required=True,
|
||||
choices=["seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"],
|
||||
help="Train type.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--data_root_path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to the root directory of the dataset.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--param_path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to the model parameters or resume checkpoint.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--resume",
|
||||
action="store_true",
|
||||
default=False,
|
||||
help="Resume training from checkpoint at --param_path "
|
||||
"(restore epoch, consumed_samples, optimizer & scheduler state).",
|
||||
)
|
||||
merged = {}
|
||||
for section in ("model", "data", "parallel", "training", "ckpt", "log"):
|
||||
if section in cfg:
|
||||
merged.update(cfg[section])
|
||||
|
||||
parser.add_argument(
|
||||
"--n_epoch", type=int, default=1, help="Number of epochs to train."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_per_device", type=int, default=1, help="Batch size per GPU."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--grad_accum_steps",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of iterations between each optimizer step.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--warmup_ratio",
|
||||
type=float,
|
||||
default=0.05,
|
||||
help="Fraction of total steps used for LR warmup.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_lr", type=float, default=3e-4, help="Max learning rate for training."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_grad_norm",
|
||||
type=float,
|
||||
default=1.0,
|
||||
help="Max gradient norm for clipping. None disables clipping.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--weight_decay",
|
||||
type=float,
|
||||
default=0.1,
|
||||
help="Weight decay (applied to Muon matrix params; non-matrix use 0).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--muon_momentum",
|
||||
type=float,
|
||||
default=0.95,
|
||||
help="Momentum factor for Muon optimizer.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--muon_nesterov",
|
||||
action=argparse.BooleanOptionalAction,
|
||||
default=True,
|
||||
help="Enable Nesterov momentum for Muon.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--muon_ns_steps",
|
||||
type=int,
|
||||
default=5,
|
||||
help="Newton-Schulz iteration steps for Muon.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--muon_adjust_lr",
|
||||
type=str,
|
||||
default="match_rms_adamw",
|
||||
choices=["original", "match_rms_adamw"],
|
||||
help="Muon learning rate adjustment strategy.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--random_seed", type=int, default=3407, help="Random seed for reproducibility."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num_workers", type=int, default=4, help="Number of workers for data loading."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--no_pin_memory",
|
||||
action="store_false",
|
||||
dest="pin_memory",
|
||||
help="Disable pin memory",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--window_size",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Max length of the input sequence.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--stride", type=int, default=None, help="Step size of the input sequence."
|
||||
)
|
||||
parser.add_argument("--dpo_beta", type=float, default=0.1, help="DPO beta value.")
|
||||
parser.add_argument("--group_size", type=int, default=4, help="GRPO group size.")
|
||||
parser.add_argument(
|
||||
"--grpo_clip_eps", type=float, default=0.2, help="GRPO clipping epsilon."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--grpo_kl_coef", type=float, default=0.01, help="GRPO KL penalty coefficient."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--label_smoothing",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="cross_entropy function label smoothing parameter",
|
||||
)
|
||||
for key, value in passed_kwargs.items():
|
||||
if value is not None:
|
||||
merged[key] = value
|
||||
|
||||
# online rollout
|
||||
parser.add_argument(
|
||||
"--rollout_interval",
|
||||
type=int,
|
||||
default=512,
|
||||
help="Number of optimizer steps between online rollouts.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--rollout_temperature",
|
||||
type=float,
|
||||
default=0.7,
|
||||
help="Sampling temperature for online rollout.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--rollout_top_k",
|
||||
type=int,
|
||||
default=0,
|
||||
help="Top-k filtering for online rollout (0=disable).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--rollout_top_p",
|
||||
type=float,
|
||||
default=0.9,
|
||||
help="Top-p (nucleus) filtering for online rollout.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--rollout_max_tokens",
|
||||
type=int,
|
||||
default=1024,
|
||||
help="Maximum generated tokens per response in rollout.",
|
||||
)
|
||||
return merged
|
||||
|
||||
parser.add_argument(
|
||||
"--gradient_checkpointing",
|
||||
action=argparse.BooleanOptionalAction,
|
||||
default=False,
|
||||
help="Enable activation checkpointing for DecoderBlock modules.",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--ckpt_interval",
|
||||
type=int,
|
||||
default=5000,
|
||||
help="Number of iters between checkpoints.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ckpt_dir",
|
||||
type=str,
|
||||
default="checkpoint",
|
||||
help="Directory to save checkpoints.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--val_split",
|
||||
type=float,
|
||||
default=None,
|
||||
help="Ratio to split from training dataset for validation (e.g. 0.05).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--val_step",
|
||||
type=int,
|
||||
default=1000,
|
||||
help="Number of optimizer steps between validation runs.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--metrics",
|
||||
nargs="*",
|
||||
default=["loss", "lr", "grad_norm"],
|
||||
help="Metrics to log (e.g. --metrics loss lr val_loss). Default: loss lr grad_norm.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--log_dir",
|
||||
type=str,
|
||||
default="checkpoint/logs",
|
||||
help="Directory for metric logs.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--start_epoch", type=int, default=0, help="Start epoch for training."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--start_samples",
|
||||
type=int,
|
||||
default=0,
|
||||
help="Start samples (per rank) for training.",
|
||||
)
|
||||
_TRAIN_TYPE = ["seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"]
|
||||
_PARALLEL = ["none", "ddp", "fsdp"]
|
||||
_SCHEDULES = ["cosine", "sgdr", "wsd"]
|
||||
_BACKENDS = ["nccl", "gloo"]
|
||||
_START_METHODS = ["spawn", "fork", "forkserver"]
|
||||
|
||||
parser.add_argument(
|
||||
"--master_addr",
|
||||
type=str,
|
||||
default="localhost",
|
||||
help="Master node address for distributed training.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--master_port",
|
||||
type=str,
|
||||
default="29500",
|
||||
help="Master node port for distributed training.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--backend",
|
||||
type=str,
|
||||
default="nccl",
|
||||
help="Distributed training backend.",
|
||||
)
|
||||
parser.add_argument("--nprocs", type=int, default=1, help="Number of GPUs to use.")
|
||||
parser.add_argument(
|
||||
"--parallel_mode",
|
||||
type=str,
|
||||
default="none",
|
||||
choices=["none", "ddp", "fsdp", "fsdp2"],
|
||||
help="Parallel training strategy (none, ddp, fsdp, fsdp2).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--device_type", type=str, default="cuda", help="Device type to use."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--start_method",
|
||||
type=str,
|
||||
default="spawn",
|
||||
choices=["spawn", "fork", "forkserver"],
|
||||
help="Multiprocessing start method.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--neftune_alpha",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="NEFTune noise alpha (0=disabled, typical: 5.0).",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--schedule_type",
|
||||
type=str,
|
||||
default="cosine",
|
||||
choices=["cosine", "sgdr", "wsd"],
|
||||
help="Learning rate scheduler type.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--min_rate",
|
||||
type=float,
|
||||
default=None,
|
||||
help="Minimum LR as fraction of base LR. Uses scheduler default if not set (cosine/sgdr: 0.05, wsd: 0.0).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--cycle_length",
|
||||
type=int,
|
||||
default=None,
|
||||
help="SGDR first cycle length in steps. Defaults to total_steps - warmup_steps.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--t_mult",
|
||||
type=int,
|
||||
default=2,
|
||||
help="SGDR cycle length multiplier per restart.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--stable_steps",
|
||||
type=int,
|
||||
default=None,
|
||||
help="WSD stable plateau steps. Required when --schedule_type wsd.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--decay_steps",
|
||||
type=int,
|
||||
default=None,
|
||||
help="WSD decay steps. Defaults to total_steps - warmup_steps - stable_steps.",
|
||||
)
|
||||
@click.command(
|
||||
name="train",
|
||||
help="Start model training (pretrain / SFT / DPO / GRPO).",
|
||||
context_settings={"show_default": True},
|
||||
)
|
||||
@click.option(
|
||||
"--config",
|
||||
"-c",
|
||||
"config_path",
|
||||
type=click.Path(exists=True),
|
||||
help="YAML config file. CLI flags override YAML values.",
|
||||
)
|
||||
@click.option(
|
||||
"--train_type",
|
||||
type=click.Choice(_TRAIN_TYPE),
|
||||
required=False,
|
||||
help="Training type.",
|
||||
)
|
||||
@click.option(
|
||||
"--data_root_path",
|
||||
type=click.Path(exists=True),
|
||||
help="Root directory of the dataset.",
|
||||
)
|
||||
@click.option(
|
||||
"--param_path",
|
||||
type=click.Path(exists=True),
|
||||
help="Path to model parameters or resume checkpoint.",
|
||||
)
|
||||
@click.option("--resume", is_flag=True, default=False, help="Resume from checkpoint.")
|
||||
@click.option("--n_epoch", type=int, default=1, help="Number of epochs.")
|
||||
@click.option("--batch_per_device", type=int, default=1, help="Batch size per GPU.")
|
||||
@click.option(
|
||||
"--grad_accum_steps", type=int, default=1, help="Gradient accumulation steps."
|
||||
)
|
||||
@click.option(
|
||||
"--warmup_ratio",
|
||||
type=float,
|
||||
default=0.05,
|
||||
help="Fraction of total steps for LR warmup.",
|
||||
)
|
||||
@click.option("--max_lr", type=float, default=3e-4, help="Max learning rate.")
|
||||
@click.option(
|
||||
"--max_grad_norm", type=float, default=1.0, help="Max gradient norm for clipping."
|
||||
)
|
||||
@click.option("--weight_decay", type=float, default=0.1, help="Weight decay.")
|
||||
@click.option("--muon_momentum", type=float, default=0.95, help="Muon momentum factor.")
|
||||
@click.option("--muon_nesterov/--no-muon_nesterov", default=True, help="Muon Nesterov.")
|
||||
@click.option("--muon_ns_steps", type=int, default=5, help="Muon Newton-Schulz steps.")
|
||||
@click.option(
|
||||
"--muon_adjust_lr",
|
||||
type=click.Choice(["original", "match_rms_adamw"]),
|
||||
default="match_rms_adamw",
|
||||
help="Muon LR adjustment strategy.",
|
||||
)
|
||||
@click.option("--random_seed", type=int, default=3407, help="Random seed.")
|
||||
@click.option("--num_workers", type=int, default=4, help="DataLoader workers.")
|
||||
@click.option("--pin_memory/--no-pin_memory", default=True, help="Pin memory.")
|
||||
@click.option(
|
||||
"--window_size", type=int, default=None, help="Max input sequence length."
|
||||
)
|
||||
@click.option("--stride", type=int, default=None, help="Step size for sliding window.")
|
||||
@click.option("--dpo_beta", type=float, default=0.1, help="DPO beta.")
|
||||
@click.option("--group_size", type=int, default=4, help="GRPO group size.")
|
||||
@click.option("--grpo_clip_eps", type=float, default=0.2, help="GRPO clip epsilon.")
|
||||
@click.option(
|
||||
"--grpo_kl_coef", type=float, default=0.01, help="GRPO KL penalty coefficient."
|
||||
)
|
||||
@click.option("--label_smoothing", type=float, default=0.0, help="Label smoothing.")
|
||||
@click.option(
|
||||
"--rollout_interval", type=int, default=512, help="Steps between rollouts."
|
||||
)
|
||||
@click.option(
|
||||
"--rollout_temperature", type=float, default=0.7, help="Rollout temperature."
|
||||
)
|
||||
@click.option("--rollout_top_k", type=int, default=0, help="Rollout top-k (0=disable).")
|
||||
@click.option("--rollout_top_p", type=float, default=0.9, help="Rollout top-p.")
|
||||
@click.option(
|
||||
"--rollout_max_tokens",
|
||||
type=int,
|
||||
default=1024,
|
||||
help="Max tokens per rollout response.",
|
||||
)
|
||||
@click.option(
|
||||
"--gradient_checkpointing/--no-gradient_checkpointing",
|
||||
default=False,
|
||||
help="Enable activation checkpointing.",
|
||||
)
|
||||
@click.option(
|
||||
"--compile",
|
||||
"compile_mode",
|
||||
type=click.Choice(["default", "reduce-overhead", "max-autotune"]),
|
||||
default=None,
|
||||
help="torch.compile mode. Omit to disable.",
|
||||
)
|
||||
@click.option(
|
||||
"--ckpt_interval", type=int, default=5000, help="Steps between checkpoints."
|
||||
)
|
||||
@click.option(
|
||||
"--ckpt_dir", type=click.Path(), default="checkpoint", help="Checkpoint directory."
|
||||
)
|
||||
@click.option("--val_split", type=float, default=None, help="Validation split ratio.")
|
||||
@click.option(
|
||||
"--val_step", type=int, default=1000, help="Steps between validation runs."
|
||||
)
|
||||
@click.option(
|
||||
"--metrics",
|
||||
multiple=True,
|
||||
default=("loss", "lr", "grad_norm"),
|
||||
help="Metrics to log (repeatable).",
|
||||
)
|
||||
@click.option("--start_epoch", type=int, default=0, help="Start epoch.")
|
||||
@click.option("--start_samples", type=int, default=0, help="Start samples (per rank).")
|
||||
@click.option(
|
||||
"--master_addr", type=str, default="localhost", help="Master node address."
|
||||
)
|
||||
@click.option("--master_port", type=str, default="29500", help="Master node port.")
|
||||
@click.option(
|
||||
"--backend",
|
||||
type=click.Choice(_BACKENDS),
|
||||
default="nccl",
|
||||
help="Distributed backend.",
|
||||
)
|
||||
@click.option("--nprocs", type=int, default=1, help="Number of GPUs.")
|
||||
@click.option(
|
||||
"--parallel_mode",
|
||||
type=click.Choice(_PARALLEL),
|
||||
default="fsdp",
|
||||
help="Parallel strategy.",
|
||||
)
|
||||
@click.option("--device_type", type=str, default="cuda", help="Device type.")
|
||||
@click.option(
|
||||
"--start_method",
|
||||
type=click.Choice(_START_METHODS),
|
||||
default="spawn",
|
||||
help="Multiprocessing start method.",
|
||||
)
|
||||
@click.option("--neftune_alpha", type=float, default=0.0, help="NEFTune noise alpha.")
|
||||
@click.option(
|
||||
"--schedule_type",
|
||||
type=click.Choice(_SCHEDULES),
|
||||
default="cosine",
|
||||
help="LR scheduler.",
|
||||
)
|
||||
@click.option(
|
||||
"--min_rate", type=float, default=None, help="Minimum LR as fraction of base LR."
|
||||
)
|
||||
@click.option("--cycle_length", type=int, default=None, help="SGDR first cycle length.")
|
||||
@click.option("--t_mult", type=int, default=2, help="SGDR cycle length multiplier.")
|
||||
@click.option(
|
||||
"--stable_steps", type=int, default=None, help="WSD stable plateau steps."
|
||||
)
|
||||
@click.option("--decay_steps", type=int, default=None, help="WSD decay steps.")
|
||||
@click.option("--tp_size", type=int, default=None, help="Tensor parallelism (future).")
|
||||
@click.option(
|
||||
"--dry-run",
|
||||
is_flag=True,
|
||||
default=False,
|
||||
help="Validate config and print plan, do not train.",
|
||||
)
|
||||
@click.pass_context
|
||||
def train_command(ctx, config_path, dry_run, metrics, **kwargs):
|
||||
"""Start model training (pretrain / SFT / DPO / GRPO)."""
|
||||
if config_path:
|
||||
kwargs = _merge_yaml_into_kwargs(config_path, kwargs)
|
||||
|
||||
args = parser.parse_args()
|
||||
required = ["train_type", "data_root_path", "param_path"]
|
||||
missing = [k for k in required if kwargs.get(k) is None]
|
||||
if missing:
|
||||
raise click.UsageError(
|
||||
f"Missing required options: {', '.join(missing)}. "
|
||||
f"Use --config YAML or provide them directly."
|
||||
)
|
||||
|
||||
return args
|
||||
# Convert tuple back to list
|
||||
kwargs["metrics"] = list(metrics)
|
||||
# Remove tp_size (not yet wired)
|
||||
kwargs.pop("tp_size", None)
|
||||
|
||||
if dry_run:
|
||||
_print_dry_run(kwargs)
|
||||
return
|
||||
|
||||
train(**kwargs)
|
||||
|
||||
|
||||
def _print_dry_run(kwargs: dict) -> None:
|
||||
"""Print training plan summary."""
|
||||
rows = [
|
||||
("Train type", kwargs.get("train_type")),
|
||||
("Model path", kwargs.get("param_path")),
|
||||
("Data path", kwargs.get("data_root_path")),
|
||||
("Parallel mode", kwargs.get("parallel_mode", "none")),
|
||||
("GPUs", str(kwargs.get("nprocs", 1))),
|
||||
("Epochs", str(kwargs.get("n_epoch", 1))),
|
||||
("Batch/device", str(kwargs.get("batch_per_device", 1))),
|
||||
("Grad accum", str(kwargs.get("grad_accum_steps", 1))),
|
||||
("Max LR", str(kwargs.get("max_lr", "?"))),
|
||||
("Schedule", str(kwargs.get("schedule_type", "cosine"))),
|
||||
("Warmup ratio", str(kwargs.get("warmup_ratio", 0.05))),
|
||||
("Window size", str(kwargs.get("window_size", "config default"))),
|
||||
("Checkpoint dir", str(kwargs.get("ckpt_dir", "checkpoint"))),
|
||||
("Checkpoint interval", str(kwargs.get("ckpt_interval", 5000))),
|
||||
("Resume", str(kwargs.get("resume", False))),
|
||||
]
|
||||
max_len = max(len(k) for k, _ in rows)
|
||||
click.secho("\n=== Training Plan (dry-run) ===", fg="cyan", bold=True)
|
||||
for key, val in rows:
|
||||
click.echo(f" {key:<{max_len}s} : {val}")
|
||||
click.secho("=" * 40, fg="cyan")
|
||||
|
||||
|
||||
def create_model(config):
|
||||
@@ -438,7 +380,6 @@ def train(
|
||||
val_split: float,
|
||||
val_step: int,
|
||||
metrics: list[str],
|
||||
log_dir: str,
|
||||
max_grad_norm: float,
|
||||
random_seed: int,
|
||||
num_workers: int,
|
||||
@@ -462,19 +403,22 @@ def train(
|
||||
decay_steps: int,
|
||||
**kwargs,
|
||||
):
|
||||
assert train_type in [
|
||||
if train_type not in [
|
||||
"seq",
|
||||
"sft",
|
||||
"dpo",
|
||||
"grpo",
|
||||
"online_grpo",
|
||||
"online_dpo",
|
||||
]
|
||||
assert os.path.exists(param_path)
|
||||
if nprocs > 1 and parallel_mode == "none":
|
||||
]:
|
||||
raise ValueError(
|
||||
"--nprocs > 1 requires --parallel_mode to be 'ddp', 'fsdp', or 'fsdp2'"
|
||||
f"Invalid train_type '{train_type}'. "
|
||||
f"Must be one of: seq, sft, dpo, grpo, online_grpo, online_dpo"
|
||||
)
|
||||
if not os.path.exists(param_path):
|
||||
raise FileNotFoundError(f"Model directory not found: {param_path}")
|
||||
if nprocs > 1 and parallel_mode == "none":
|
||||
raise ValueError("--nprocs > 1 requires --parallel_mode to be 'ddp' or 'fsdp'")
|
||||
|
||||
# Load config
|
||||
config_path = os.path.join(param_path, "config.json")
|
||||
@@ -497,7 +441,7 @@ def train(
|
||||
rollout_top_k = kwargs.pop("rollout_top_k", 0)
|
||||
rollout_top_p = kwargs.pop("rollout_top_p", 0.9)
|
||||
rollout_max_tokens = kwargs.pop("rollout_max_tokens", 1024)
|
||||
reward_model_fn: Optional[Callable[[], BaseRewardModel]] = None
|
||||
reward_model_fn: Callable[[], BaseRewardModel] | None = None
|
||||
|
||||
executor_kwargs = {}
|
||||
if parallel_mode == "ddp":
|
||||
@@ -556,6 +500,7 @@ def train(
|
||||
)
|
||||
|
||||
grad_ckpt_modules = [DecoderBlock] if gradient_checkpointing else []
|
||||
compile_mode = kwargs.pop("compile_mode", None)
|
||||
|
||||
collate_fn = None
|
||||
if train_type == "dpo":
|
||||
@@ -592,8 +537,8 @@ def train(
|
||||
val_split=val_split,
|
||||
val_step=val_step,
|
||||
metrics=metrics,
|
||||
log_dir=log_dir,
|
||||
gradient_checkpointing_modules=grad_ckpt_modules,
|
||||
compile_mode=compile_mode,
|
||||
executor_kwargs=executor_kwargs,
|
||||
extra_kwargs=strategy_kwargs,
|
||||
neftune_alpha=neftune_alpha,
|
||||
@@ -611,5 +556,5 @@ def train(
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = parse_args()
|
||||
train(**vars(args))
|
||||
setup_logging()
|
||||
train_command()
|
||||
|
||||
+21
-106
@@ -6,11 +6,16 @@ import tempfile
|
||||
import pytest
|
||||
import torch
|
||||
from tokenizers import Tokenizer, models, pre_tokenizers, trainers
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.extension import KERNEL_NAMES, is_available
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
from tests.helpers import TINY_CONFIG, RandomTokenDataset, make_tiny_config
|
||||
|
||||
CUDA_AVAIL = torch.cuda.is_available()
|
||||
KERNEL_AVAIL = CUDA_AVAIL and all(is_available(k) for k in KERNEL_NAMES)
|
||||
skip_no_cuda = pytest.mark.skipif(not CUDA_AVAIL, reason="CUDA not available")
|
||||
skip_no_kernel = pytest.mark.skipif(not KERNEL_AVAIL, reason="CUDA kernels not built")
|
||||
|
||||
|
||||
def pytest_configure(config):
|
||||
@@ -19,6 +24,12 @@ def pytest_configure(config):
|
||||
config.addinivalue_line("markers", "unit: fast unit tests")
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def device():
|
||||
"""Session-scoped device string (``"cuda"`` if available, else ``"cpu"``)."""
|
||||
return "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
|
||||
def create_test_tokenizer(vocab_size: int = 1000) -> AutoTokenizer:
|
||||
"""Create a simple tokenizer for testing purposes."""
|
||||
tokenizer = Tokenizer(models.BPE())
|
||||
@@ -33,69 +44,6 @@ def create_test_tokenizer(vocab_size: int = 1000) -> AutoTokenizer:
|
||||
return auto_tokenizer
|
||||
|
||||
|
||||
class RandomDataset(Dataset):
|
||||
"""Random dataset for testing purposes."""
|
||||
|
||||
def __init__(self, length=None, max_length=64, vocab_size=1000):
|
||||
self.length = length or int(torch.randint(100, 200, (1,)).item())
|
||||
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,)),
|
||||
}
|
||||
|
||||
|
||||
class MultiTurnDataset(Dataset):
|
||||
"""Multi-turn dataset with loss mask for SFT training tests."""
|
||||
|
||||
def __init__(self, length=None, max_length=64, vocab_size=1000):
|
||||
self.length = length or int(torch.randint(100, 200, (1,)).item())
|
||||
self.max_length = max_length
|
||||
self.vocab_size = vocab_size
|
||||
|
||||
def __len__(self):
|
||||
return self.length
|
||||
|
||||
def __getitem__(self, idx):
|
||||
input_ids = torch.randint(0, self.vocab_size, (self.max_length,))
|
||||
target_ids = torch.randint(0, self.vocab_size, (self.max_length,))
|
||||
loss_mask = torch.randint(0, 1, (self.max_length,))
|
||||
|
||||
return {
|
||||
"input_ids": input_ids,
|
||||
"target_ids": target_ids,
|
||||
"loss_mask": loss_mask,
|
||||
}
|
||||
|
||||
|
||||
class EarlyStoppingDataset(Dataset):
|
||||
"""Dataset that triggers early stopping after consuming a specified number of samples."""
|
||||
|
||||
def __init__(self, length=10, stop_after=5):
|
||||
self.length = length
|
||||
self.stop_after = stop_after
|
||||
self.count = 0
|
||||
|
||||
def __len__(self):
|
||||
return self.length
|
||||
|
||||
def __getitem__(self, idx):
|
||||
self.count += 1
|
||||
if self.count == self.stop_after:
|
||||
raise RuntimeError("Simulated early stopping")
|
||||
|
||||
return {
|
||||
"input_ids": torch.randint(0, 1000, (64,)),
|
||||
"target_ids": torch.randint(0, 1000, (64,)),
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def test_tokenizer():
|
||||
"""Session-scoped tokenizer, created once for the entire test run."""
|
||||
@@ -103,50 +51,20 @@ def test_tokenizer():
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def test_model():
|
||||
def test_model(device):
|
||||
"""Session-scoped small AutoRegressiveLM model, created once."""
|
||||
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"
|
||||
config = make_tiny_config()
|
||||
model = AutoRegressiveLM(config).to(device=device)
|
||||
|
||||
return {
|
||||
"model": model,
|
||||
"device": device,
|
||||
"config": config,
|
||||
}
|
||||
return {"model": model, "device": device, "config": config}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def base_test_env(test_model, test_tokenizer):
|
||||
"""Function-scoped test environment with isolated temp directory.
|
||||
|
||||
Composes session-scoped model and tokenizer with a per-test temp dir.
|
||||
"""
|
||||
"""Function-scoped test environment with isolated temp directory."""
|
||||
test_dir = tempfile.mkdtemp()
|
||||
config_path = os.path.join(test_dir, "config.json")
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(
|
||||
{
|
||||
"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,
|
||||
},
|
||||
f,
|
||||
)
|
||||
json.dump(TINY_CONFIG, f)
|
||||
|
||||
yield {
|
||||
"device": test_model["device"],
|
||||
@@ -162,17 +80,14 @@ def base_test_env(test_model, test_tokenizer):
|
||||
|
||||
@pytest.fixture
|
||||
def random_dataset():
|
||||
dataset = RandomDataset()
|
||||
yield dataset
|
||||
return RandomTokenDataset(length=None)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def multi_turn_dataset():
|
||||
dataset = MultiTurnDataset()
|
||||
yield dataset
|
||||
return RandomTokenDataset(length=None, with_loss_mask=True)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def early_stopping_dataset():
|
||||
dataset = EarlyStoppingDataset()
|
||||
yield dataset
|
||||
return RandomTokenDataset(length=10, stop_after=5)
|
||||
|
||||
+42
-46
@@ -7,18 +7,24 @@ import pytest
|
||||
import torch
|
||||
|
||||
from astrai.config.preprocess_config import PipelineConfig
|
||||
from astrai.dataset.dataset import DatasetFactory, dpo_tokenize
|
||||
from astrai.dataset.dataset import (
|
||||
DatasetFactory,
|
||||
GRPODataset,
|
||||
dpo_tokenize,
|
||||
grpo_collate_fn,
|
||||
)
|
||||
from astrai.dataset.storage import (
|
||||
H5Store,
|
||||
JsonlStore,
|
||||
MmapStore,
|
||||
StoreFactory,
|
||||
detect_format,
|
||||
)
|
||||
from astrai.preprocessing.builder import SectionedMaskBuilder
|
||||
from astrai.serialization import (
|
||||
load_bin,
|
||||
save_bin,
|
||||
save_h5,
|
||||
)
|
||||
from tests.data.conftest import make_grpo_no_template_config
|
||||
|
||||
|
||||
def _rand_seq(length, vocab=1000):
|
||||
@@ -70,7 +76,7 @@ def _make_seq_dataset(
|
||||
):
|
||||
if data is None:
|
||||
data = {"sequence": [_rand_seq(seq_length)]}
|
||||
save_h5(test_dir, name, data)
|
||||
save_bin(test_dir, data)
|
||||
return DatasetFactory.load(
|
||||
train_type,
|
||||
test_dir,
|
||||
@@ -83,12 +89,15 @@ def test_dataset_loader_random_paths(base_test_env):
|
||||
"""Test dataset loader with multiple random paths"""
|
||||
test_dir = base_test_env["test_dir"]
|
||||
|
||||
loaded_dataset = None
|
||||
num_files = np.random.randint(2, 5)
|
||||
for i in range(num_files):
|
||||
seq_length = np.random.randint(200, 400)
|
||||
dummy_data = {"sequence": [_rand_seq(seq_length) for _ in range(10)]}
|
||||
sub_dir = os.path.join(test_dir, f"sub_{i}")
|
||||
os.makedirs(sub_dir, exist_ok=True)
|
||||
loaded_dataset = _make_seq_dataset(
|
||||
test_dir, f"data_{i}", seq_length, data=dummy_data
|
||||
sub_dir, f"data_{i}", seq_length, data=dummy_data
|
||||
)
|
||||
assert loaded_dataset is not None
|
||||
assert len(loaded_dataset) > 0
|
||||
@@ -113,8 +122,15 @@ def test_dpo_strategy_with_random_data(base_test_env):
|
||||
"chosen_mask": [torch.ones(seq_length, dtype=torch.bool)],
|
||||
"rejected_mask": [torch.ones(seq_length, dtype=torch.bool)],
|
||||
}
|
||||
dpo_dataset = _make_seq_dataset(
|
||||
test_dir, "dpo_data", seq_length, train_type="dpo", data=dummy_data
|
||||
save_bin(
|
||||
test_dir,
|
||||
dummy_data,
|
||||
record_keys=["chosen", "rejected", "chosen_mask", "rejected_mask"],
|
||||
)
|
||||
dpo_dataset = DatasetFactory.load(
|
||||
train_type="dpo",
|
||||
load_path=test_dir,
|
||||
window_size=0,
|
||||
)
|
||||
|
||||
assert dpo_dataset is not None
|
||||
@@ -196,24 +212,20 @@ def test_dataset_too_short_for_window(base_test_env):
|
||||
|
||||
def test_unloaded_sample_window_raises():
|
||||
"""Store.sample_window before load raises RuntimeError."""
|
||||
from astrai.dataset.storage import H5Store
|
||||
|
||||
store = H5Store(window_size=64, stride=64)
|
||||
store = MmapStore(window_size=64, stride=64)
|
||||
with pytest.raises(IndexError, match="Data too short"):
|
||||
store.sample_window(0)
|
||||
|
||||
|
||||
def test_unloaded_dataset_len():
|
||||
"""__len__ on a store with no data returns 0."""
|
||||
from astrai.dataset.storage import H5Store
|
||||
|
||||
store = H5Store(window_size=64, stride=64)
|
||||
store = MmapStore(window_size=64, stride=64)
|
||||
assert len(store) == 0
|
||||
|
||||
|
||||
def test_store_unloaded_len():
|
||||
"""Unloaded Store has __len__ == 0"""
|
||||
store = H5Store()
|
||||
store = MmapStore()
|
||||
assert len(store) == 0
|
||||
assert store.keys == []
|
||||
|
||||
@@ -227,7 +239,7 @@ def test_store_fetch_begin_equals_end(base_test_env):
|
||||
|
||||
def test_store_fetch_before_load():
|
||||
"""Store.fetch before load raises RuntimeError"""
|
||||
store = H5Store()
|
||||
store = MmapStore()
|
||||
with pytest.raises(RuntimeError, match="not loaded"):
|
||||
store.fetch(0, 10, "sequence")
|
||||
|
||||
@@ -255,9 +267,7 @@ def test_create_store_invalid_type():
|
||||
|
||||
|
||||
def test_store_multi_segment_concat(base_test_env):
|
||||
"""Multi-segment H5 data is concatenated into single tensor at load time"""
|
||||
import os
|
||||
|
||||
"""Multi-segment data is concatenated into single tensor at load time"""
|
||||
test_dir = base_test_env["test_dir"]
|
||||
data_dir = os.path.join(test_dir, "multi_seg")
|
||||
os.makedirs(data_dir, exist_ok=True)
|
||||
@@ -267,9 +277,9 @@ def test_store_multi_segment_concat(base_test_env):
|
||||
torch.tensor([4, 5, 6, 7]),
|
||||
torch.tensor([8, 9]),
|
||||
]
|
||||
save_h5(data_dir, "data", {"sequence": segs})
|
||||
save_bin(data_dir, {"sequence": segs})
|
||||
|
||||
store = StoreFactory.create("h5")
|
||||
store = StoreFactory.create("bin")
|
||||
store.load(data_dir)
|
||||
assert store.token_count == 9
|
||||
result = store.fetch(2, 7, "sequence")
|
||||
@@ -321,7 +331,7 @@ def test_mmap_dataset_load(base_test_env):
|
||||
|
||||
def test_normalize_empty_key():
|
||||
"""_normalize with empty tensor list does not crash."""
|
||||
store = H5Store()
|
||||
store = MmapStore()
|
||||
store._normalize({"sequence": []})
|
||||
assert len(store) == 0
|
||||
assert store.num_records == 0 # empty key forces num_records=0
|
||||
@@ -330,7 +340,7 @@ def test_normalize_empty_key():
|
||||
|
||||
def test_normalize_mixed_empty_key():
|
||||
"""_normalize with empty + non-empty keys returns min=0 records."""
|
||||
store = H5Store()
|
||||
store = MmapStore()
|
||||
store._normalize({"sequence": [torch.tensor([1, 2, 3])], "loss_mask": []})
|
||||
assert len(store) == 0
|
||||
assert store.num_records == 0
|
||||
@@ -340,8 +350,6 @@ def test_normalize_mixed_empty_key():
|
||||
|
||||
def test_grpo_dataset_dtype(base_test_env):
|
||||
"""GRPO dataset returns correct dtypes for per-record structured data."""
|
||||
from astrai.dataset.dataset import GRPODataset
|
||||
|
||||
G = 4
|
||||
store = type(
|
||||
"FakeStore",
|
||||
@@ -373,8 +381,6 @@ def test_grpo_dataset_dtype(base_test_env):
|
||||
|
||||
def test_grpo_dataset_load(base_test_env):
|
||||
"""GRPO dataset loads record-structured data with per-response boundaries."""
|
||||
from astrai.dataset.dataset import GRPODataset
|
||||
|
||||
G = 3
|
||||
prompt_len = 8
|
||||
resp_lens = [5, 7, 4]
|
||||
@@ -430,15 +436,14 @@ def test_detect_format_bin_dir(base_test_env):
|
||||
|
||||
def test_store_fetch_multi_key(base_test_env):
|
||||
test_dir = base_test_env["test_dir"]
|
||||
save_h5(
|
||||
save_bin(
|
||||
test_dir,
|
||||
"multi_key",
|
||||
{
|
||||
"sequence": [torch.randint(0, 100, (100,), dtype=torch.int64)],
|
||||
"loss_mask": [torch.ones(100, dtype=torch.int64)],
|
||||
},
|
||||
)
|
||||
store = StoreFactory.create("h5")
|
||||
store = StoreFactory.create("bin")
|
||||
store.load(test_dir)
|
||||
result = store.fetch(10, 20, ["sequence", "loss_mask"])
|
||||
assert isinstance(result, dict)
|
||||
@@ -448,8 +453,8 @@ def test_store_fetch_multi_key(base_test_env):
|
||||
|
||||
def test_store_fetch_out_of_bounds(base_test_env):
|
||||
test_dir = base_test_env["test_dir"]
|
||||
save_h5(test_dir, "bounds", {"sequence": [torch.randint(0, 100, (50,))]})
|
||||
store = StoreFactory.create("h5")
|
||||
save_bin(test_dir, {"sequence": [torch.randint(0, 100, (50,))]})
|
||||
store = StoreFactory.create("bin")
|
||||
store.load(test_dir)
|
||||
with pytest.raises(ValueError, match="out of bounds"):
|
||||
store.fetch(-1, 10, "sequence")
|
||||
@@ -461,7 +466,7 @@ def test_store_fetch_out_of_bounds(base_test_env):
|
||||
|
||||
def test_dataset_load_explicit_storage_type(base_test_env):
|
||||
test_dir = base_test_env["test_dir"]
|
||||
dataset = _make_seq_dataset(test_dir, "explicit", storage_type="h5")
|
||||
dataset = _make_seq_dataset(test_dir, "explicit", storage_type="bin")
|
||||
assert len(dataset) > 0
|
||||
assert dataset.token_count == 200
|
||||
|
||||
@@ -820,9 +825,6 @@ def _write_grpo_jsonl(test_dir, tokenizer_path, records):
|
||||
|
||||
def test_grpo_builder_preserves_response_boundaries(base_test_env):
|
||||
"""MultiOutputMaskBuilder with list_field returns List[List[int]] for responses."""
|
||||
from astrai.preprocessing.builder import SectionedMaskBuilder
|
||||
from tests.data.conftest import make_grpo_no_template_config
|
||||
|
||||
tokenizer = base_test_env["tokenizer"]
|
||||
_save_test_tokenizer(base_test_env["test_dir"], tokenizer)
|
||||
|
||||
@@ -863,8 +865,6 @@ def test_grpo_builder_preserves_response_boundaries(base_test_env):
|
||||
|
||||
def test_grpo_end_to_end_jsonl(base_test_env):
|
||||
"""Full GRPO pipeline: JSONL → JsonlStore → GRPODataset → collate_fn."""
|
||||
from astrai.dataset.dataset import grpo_collate_fn
|
||||
|
||||
test_dir = base_test_env["test_dir"]
|
||||
tokenizer = base_test_env["tokenizer"]
|
||||
tokenizer_path = _save_test_tokenizer(test_dir, tokenizer)
|
||||
@@ -913,8 +913,6 @@ def test_grpo_end_to_end_jsonl(base_test_env):
|
||||
|
||||
def test_grpo_collate_variable_lengths():
|
||||
"""collate_fn pads variable-length responses to [B, G, R_max]."""
|
||||
from astrai.dataset.dataset import grpo_collate_fn
|
||||
|
||||
batch = [
|
||||
{
|
||||
"prompts": torch.tensor([1, 2, 3]),
|
||||
@@ -954,8 +952,6 @@ def test_grpo_collate_variable_lengths():
|
||||
|
||||
def test_grpo_multiple_records(base_test_env):
|
||||
"""GRPODataset loads multiple records with correct structure."""
|
||||
from astrai.dataset.dataset import GRPODataset
|
||||
|
||||
G = 4
|
||||
n_records = 5
|
||||
|
||||
@@ -1128,8 +1124,8 @@ def test_jsonl_store_eager_len_returns_token_count(base_test_env):
|
||||
assert len(store.keys) > 0
|
||||
|
||||
|
||||
def test_h5_store_dual_mode(base_test_env):
|
||||
"""H5Store supports both fetch (stream) and fetch_record (record).
|
||||
def test_mmap_store_dual_mode(base_test_env):
|
||||
"""MmapStore supports both fetch (stream) and fetch_record (record).
|
||||
|
||||
No window configured → ``len(store)`` reflects the record count
|
||||
(2). ``token_count`` retains the legacy stream length (128), and
|
||||
@@ -1143,9 +1139,9 @@ def test_h5_store_dual_mode(base_test_env):
|
||||
"chosen": [_rand_seq(seq_length), _rand_seq(seq_length)],
|
||||
"rejected": [_rand_seq(seq_length), _rand_seq(seq_length)],
|
||||
}
|
||||
save_h5(test_dir, "dpo_data", dummy_data)
|
||||
save_bin(test_dir, dummy_data, record_keys=["chosen", "rejected"])
|
||||
|
||||
store = H5Store()
|
||||
store = MmapStore()
|
||||
store.load(test_dir)
|
||||
|
||||
assert store.token_count == seq_length * 2
|
||||
@@ -1160,7 +1156,7 @@ def test_h5_store_dual_mode(base_test_env):
|
||||
|
||||
# Window-configured view of the same data uses stream sample count:
|
||||
# token_count=128, window_size=64 → num_samples = (128-1-64)//64 + 1 = 1
|
||||
stream_view = H5Store(window_size=seq_length, stride=seq_length)
|
||||
stream_view = MmapStore(window_size=seq_length, stride=seq_length)
|
||||
stream_view.load(test_dir)
|
||||
assert len(stream_view) == 1
|
||||
|
||||
|
||||
@@ -30,7 +30,7 @@ def test_from_dict_flat():
|
||||
"mask": {"system": "mask", "assistant": "train"},
|
||||
"mask_default": "mask",
|
||||
"preprocessing": {"max_seq_len": 1024},
|
||||
"output": {"storage_format": "h5"},
|
||||
"output": {"storage_format": "bin"},
|
||||
}
|
||||
config = PipelineConfig.from_dict(data)
|
||||
assert config.input.sections == [
|
||||
@@ -38,7 +38,7 @@ def test_from_dict_flat():
|
||||
]
|
||||
assert config.mask == {"system": "mask", "assistant": "train"}
|
||||
assert config.preprocessing.max_seq_len == 1024
|
||||
assert config.output.storage_format == "h5"
|
||||
assert config.output.storage_format == "bin"
|
||||
|
||||
|
||||
def test_to_dict_roundtrip():
|
||||
|
||||
@@ -16,6 +16,7 @@ from tests.data.conftest import (
|
||||
make_dpo_chat_config,
|
||||
make_grpo_no_template_config,
|
||||
)
|
||||
from tests.helpers import load_shard_meta
|
||||
|
||||
|
||||
def test_filter_by_length():
|
||||
@@ -68,10 +69,7 @@ def test_full_chat_pipeline(temp_dir, chat_tokenizer_dir):
|
||||
tokenizer_path=chat_tokenizer_dir,
|
||||
).run()
|
||||
|
||||
meta_path = os.path.join(out_dir, "__default__", "shard_0000", "meta.json")
|
||||
assert os.path.exists(meta_path)
|
||||
with open(meta_path, "r") as f:
|
||||
meta = json.load(f)
|
||||
meta = load_shard_meta(out_dir)
|
||||
assert "sequence" in meta
|
||||
assert "loss_mask" in meta
|
||||
assert meta["sequence"]["dtype"] == "int32"
|
||||
@@ -112,10 +110,7 @@ def test_full_text_pipeline(temp_dir, tokenizer_dir):
|
||||
tokenizer_path=tokenizer_dir,
|
||||
).run()
|
||||
|
||||
meta_path = os.path.join(out_dir, "__default__", "shard_0000", "meta.json")
|
||||
assert os.path.exists(meta_path)
|
||||
with open(meta_path, "r") as f:
|
||||
meta = json.load(f)
|
||||
meta = load_shard_meta(out_dir)
|
||||
assert "sequence" in meta
|
||||
assert "loss_mask" not in meta
|
||||
|
||||
@@ -158,10 +153,7 @@ def test_full_instruction_pipeline(temp_dir, tokenizer_dir):
|
||||
tokenizer_path=tokenizer_dir,
|
||||
).run()
|
||||
|
||||
meta_path = os.path.join(out_dir, "__default__", "shard_0000", "meta.json")
|
||||
assert os.path.exists(meta_path)
|
||||
with open(meta_path, "r") as f:
|
||||
meta = json.load(f)
|
||||
meta = load_shard_meta(out_dir)
|
||||
assert "sequence" in meta
|
||||
assert "loss_mask" in meta
|
||||
|
||||
@@ -187,9 +179,7 @@ def test_dtype_override(temp_dir, tokenizer_dir):
|
||||
tokenizer_path=tokenizer_dir,
|
||||
).run()
|
||||
|
||||
meta_path = os.path.join(out_dir, "__default__", "shard_0000", "meta.json")
|
||||
with open(meta_path, "r") as f:
|
||||
meta = json.load(f)
|
||||
meta = load_shard_meta(out_dir)
|
||||
assert meta["sequence"]["dtype"] == "int32"
|
||||
assert meta["loss_mask"]["dtype"] == "bool"
|
||||
|
||||
@@ -221,10 +211,7 @@ def test_dpo_pipeline(temp_dir, chat_tokenizer_dir):
|
||||
tokenizer_path=chat_tokenizer_dir,
|
||||
).run()
|
||||
|
||||
meta_path = os.path.join(out_dir, "__default__", "shard_0000", "meta.json")
|
||||
assert os.path.exists(meta_path)
|
||||
with open(meta_path, "r") as f:
|
||||
meta = json.load(f)
|
||||
meta = load_shard_meta(out_dir)
|
||||
assert "chosen" in meta
|
||||
assert "rejected" in meta
|
||||
assert "chosen_mask" in meta
|
||||
@@ -254,10 +241,7 @@ def test_grpo_pipeline(temp_dir, tokenizer_dir):
|
||||
tokenizer_path=tokenizer_dir,
|
||||
).run()
|
||||
|
||||
meta_path = os.path.join(out_dir, "__default__", "shard_0000", "meta.json")
|
||||
assert os.path.exists(meta_path)
|
||||
with open(meta_path, "r") as f:
|
||||
meta = json.load(f)
|
||||
meta = load_shard_meta(out_dir)
|
||||
assert "prompts" in meta
|
||||
assert "responses" in meta
|
||||
assert "masks" in meta
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
"""Shared fixtures for extension tests."""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
from tests.conftest import skip_no_kernel
|
||||
|
||||
D = 64
|
||||
CFG = dict(
|
||||
vocab_size=1000,
|
||||
hidden_size=128,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=1,
|
||||
intermediate_size=256,
|
||||
max_position_embeddings=64,
|
||||
num_hidden_layers=2,
|
||||
rms_norm_eps=1e-5,
|
||||
attn_type="gqa",
|
||||
ffn_type="mlp",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def cuda_model():
|
||||
config = AutoRegressiveLMConfig(**CFG)
|
||||
model = AutoRegressiveLM(config).to(device="cuda", dtype=torch.bfloat16)
|
||||
model.eval()
|
||||
return model, config
|
||||
@@ -0,0 +1,45 @@
|
||||
"""Backend selection and context-manager switching tests.
|
||||
|
||||
These tests do not require CUDA — they only check that the active
|
||||
backend is correctly set and restored.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from astrai.extension import (
|
||||
ATTN_BACKEND,
|
||||
CudaBackend,
|
||||
TorchNativeBackend,
|
||||
attn_backend,
|
||||
get_backend,
|
||||
)
|
||||
|
||||
|
||||
def test_default_backend_is_torch_native():
|
||||
backend = get_backend()
|
||||
assert isinstance(backend, TorchNativeBackend)
|
||||
|
||||
|
||||
def test_attn_backend_context_with_enum():
|
||||
with attn_backend(ATTN_BACKEND.CUDA):
|
||||
assert isinstance(get_backend(), CudaBackend)
|
||||
assert isinstance(get_backend(), TorchNativeBackend)
|
||||
|
||||
|
||||
def test_attn_backend_context_with_class():
|
||||
with attn_backend(CudaBackend):
|
||||
assert isinstance(get_backend(), CudaBackend)
|
||||
assert isinstance(get_backend(), TorchNativeBackend)
|
||||
|
||||
|
||||
def test_attn_backend_context_with_instance():
|
||||
custom = CudaBackend()
|
||||
with attn_backend(custom):
|
||||
assert get_backend() is custom
|
||||
assert isinstance(get_backend(), TorchNativeBackend)
|
||||
|
||||
|
||||
def test_cudabackend_is_context_manager():
|
||||
with CudaBackend():
|
||||
assert isinstance(get_backend(), CudaBackend)
|
||||
assert isinstance(get_backend(), TorchNativeBackend)
|
||||
@@ -0,0 +1,199 @@
|
||||
"""Numerical equivalence between TorchNativeBackend and CudaBackend.
|
||||
|
||||
Covers training forward, inference prefill, inference decode (mixed
|
||||
seq_lens with padding mask), and end-to-end scheduler.run_batch.
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.extension import ATTN_BACKEND, attn_backend
|
||||
from astrai.inference.core.cache import PagePool
|
||||
from tests.extension.conftest import D, skip_no_kernel
|
||||
|
||||
|
||||
@skip_no_kernel
|
||||
def test_training_forward_matches_torch(cuda_model):
|
||||
"""Training forward (kv_cache=None) should produce identical logits."""
|
||||
model, _ = cuda_model
|
||||
input_ids = torch.randint(0, 1000, (2, 16), device="cuda")
|
||||
|
||||
with torch.no_grad():
|
||||
out_torch = model(input_ids)
|
||||
with attn_backend(ATTN_BACKEND.CUDA):
|
||||
with torch.no_grad():
|
||||
out_cuda = model(input_ids)
|
||||
|
||||
diff = (out_torch["logits"].float() - out_cuda["logits"].float()).abs().max().item()
|
||||
assert diff == 0.0, f"Training forward diff {diff} should be 0"
|
||||
|
||||
|
||||
@skip_no_kernel
|
||||
def test_prefill_with_kv_cache_matches_torch(cuda_model):
|
||||
"""Inference prefill with KV cache should match torch backend."""
|
||||
model, _ = cuda_model
|
||||
prompt_ids = [[1, 2, 3, 4, 5, 6, 7, 8], [10, 11, 12, 13, 14, 15]]
|
||||
max_len = max(len(p) for p in prompt_ids)
|
||||
batch = len(prompt_ids)
|
||||
|
||||
device = "cuda"
|
||||
input_ids = torch.zeros(batch, max_len, dtype=torch.long, device=device)
|
||||
input_mask = torch.zeros(batch, max_len, dtype=torch.bool, device=device)
|
||||
position_ids = torch.zeros(batch, max_len, dtype=torch.long, device=device)
|
||||
for i, p in enumerate(prompt_ids):
|
||||
input_ids[i, : len(p)] = torch.tensor(p, device=device)
|
||||
input_mask[i, : len(p)] = True
|
||||
position_ids[i, : len(p)] = torch.arange(len(p), device=device)
|
||||
|
||||
cache = PagePool(
|
||||
n_layers=2,
|
||||
n_kv_heads=1,
|
||||
head_dim=D,
|
||||
max_batch_size=4,
|
||||
max_seq_len=64,
|
||||
device=device,
|
||||
dtype=torch.bfloat16,
|
||||
)
|
||||
|
||||
cache.task_alloc("t1", prompt_ids[0])
|
||||
cache.task_alloc("t2", prompt_ids[1])
|
||||
kv1 = cache.bind_tasks(
|
||||
["t1", "t2"], [len(prompt_ids[0]), len(prompt_ids[1])], device, start_pos=0
|
||||
)
|
||||
with torch.inference_mode():
|
||||
out_torch = model(
|
||||
input_ids, input_mask=input_mask, kv_cache=kv1, position_ids=position_ids
|
||||
)
|
||||
|
||||
cache.task_free("t1")
|
||||
cache.task_free("t2")
|
||||
cache.task_alloc("t1", prompt_ids[0])
|
||||
cache.task_alloc("t2", prompt_ids[1])
|
||||
kv2 = cache.bind_tasks(
|
||||
["t1", "t2"], [len(prompt_ids[0]), len(prompt_ids[1])], device, start_pos=0
|
||||
)
|
||||
with attn_backend(ATTN_BACKEND.CUDA):
|
||||
with torch.inference_mode():
|
||||
out_cuda = model(
|
||||
input_ids,
|
||||
input_mask=input_mask,
|
||||
kv_cache=kv2,
|
||||
position_ids=position_ids,
|
||||
)
|
||||
|
||||
for i, p in enumerate(prompt_ids):
|
||||
d = (
|
||||
(
|
||||
out_torch["logits"][i, : len(p)].float()
|
||||
- out_cuda["logits"][i, : len(p)].float()
|
||||
)
|
||||
.abs()
|
||||
.max()
|
||||
.item()
|
||||
)
|
||||
assert d == 0.0, f"Prefill diff for sample {i}: {d}"
|
||||
|
||||
|
||||
@skip_no_kernel
|
||||
def test_decode_mixed_seq_lens_matches_torch(cuda_model):
|
||||
"""Decode with mixed seq_lens in batch — padding mask must produce correct output."""
|
||||
model, _ = cuda_model
|
||||
device = "cuda"
|
||||
|
||||
prompt_ids = [[1, 2, 3, 4, 5, 6, 7, 8], [10, 11, 12, 13, 14, 15]]
|
||||
cache = PagePool(
|
||||
n_layers=2,
|
||||
n_kv_heads=1,
|
||||
head_dim=D,
|
||||
max_batch_size=4,
|
||||
max_seq_len=64,
|
||||
device=device,
|
||||
dtype=torch.bfloat16,
|
||||
)
|
||||
|
||||
# Prefill to populate cache
|
||||
max_len = max(len(p) for p in prompt_ids)
|
||||
batch = len(prompt_ids)
|
||||
input_ids = torch.zeros(batch, max_len, dtype=torch.long, device=device)
|
||||
input_mask = torch.zeros(batch, max_len, dtype=torch.bool, device=device)
|
||||
position_ids = torch.zeros(batch, max_len, dtype=torch.long, device=device)
|
||||
for i, p in enumerate(prompt_ids):
|
||||
input_ids[i, : len(p)] = torch.tensor(p, device=device)
|
||||
input_mask[i, : len(p)] = True
|
||||
position_ids[i, : len(p)] = torch.arange(len(p), device=device)
|
||||
|
||||
cache.task_alloc("t1", prompt_ids[0])
|
||||
cache.task_alloc("t2", prompt_ids[1])
|
||||
kv = cache.bind_tasks(
|
||||
["t1", "t2"], [len(prompt_ids[0]), len(prompt_ids[1])], device, start_pos=0
|
||||
)
|
||||
with torch.inference_mode():
|
||||
model(input_ids, input_mask=input_mask, kv_cache=kv, position_ids=position_ids)
|
||||
|
||||
# Decode step — seq_lens are 9 and 7 (after extending)
|
||||
dec_ids = torch.tensor([[99], [98]], dtype=torch.long, device=device)
|
||||
dec_pos = torch.tensor([[8], [6]], dtype=torch.long, device=device)
|
||||
total_len = 9
|
||||
dec_mask = dec_pos[:, None, None] >= torch.arange(total_len, device=device)
|
||||
|
||||
kv_t = cache.bind_tasks(["t1", "t2"], [9, 7], device)
|
||||
with torch.inference_mode():
|
||||
out_torch = model(
|
||||
dec_ids, input_mask=dec_mask, kv_cache=kv_t, position_ids=dec_pos
|
||||
)
|
||||
|
||||
kv_c = cache.bind_tasks(["t1", "t2"], [9, 7], device)
|
||||
with attn_backend(ATTN_BACKEND.CUDA):
|
||||
with torch.inference_mode():
|
||||
out_cuda = model(
|
||||
dec_ids, input_mask=dec_mask, kv_cache=kv_c, position_ids=dec_pos
|
||||
)
|
||||
|
||||
diff = (out_torch["logits"].float() - out_cuda["logits"].float()).abs().max().item()
|
||||
assert diff < 0.05, f"Decode diff (mixed seq_lens): {diff}"
|
||||
|
||||
|
||||
@skip_no_kernel
|
||||
def test_run_batch_cuda_matches_torch_greedy(cuda_model):
|
||||
"""Greedy decode (temperature=0) should produce identical tokens."""
|
||||
from astrai.inference.core.scheduler import InferenceScheduler
|
||||
from tests.helpers import FakeTokenizer
|
||||
|
||||
model, _ = cuda_model
|
||||
tokenizer = FakeTokenizer()
|
||||
|
||||
prompts = [[1, 2, 3, 4, 5], [10, 11, 12, 13, 14, 15, 16]]
|
||||
|
||||
sched = InferenceScheduler(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=4,
|
||||
max_seq_len=64,
|
||||
device="cuda",
|
||||
dtype=torch.bfloat16,
|
||||
)
|
||||
out_torch = sched.run_batch(prompts, max_tokens=5, temperature=0.0)
|
||||
sched.stop()
|
||||
|
||||
cache_cuda = PagePool(
|
||||
n_layers=2,
|
||||
n_kv_heads=1,
|
||||
head_dim=D,
|
||||
max_batch_size=4,
|
||||
max_seq_len=64,
|
||||
device="cuda",
|
||||
dtype=torch.bfloat16,
|
||||
)
|
||||
sched2 = InferenceScheduler(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=4,
|
||||
max_seq_len=64,
|
||||
device="cuda",
|
||||
dtype=torch.bfloat16,
|
||||
cache=cache_cuda,
|
||||
)
|
||||
with attn_backend(ATTN_BACKEND.CUDA):
|
||||
out_cuda = sched2.run_batch(prompts, max_tokens=5, temperature=0.0)
|
||||
sched2.stop()
|
||||
|
||||
assert out_torch == out_cuda, f"Torch={out_torch} != CUDA={out_cuda}"
|
||||
@@ -0,0 +1,74 @@
|
||||
"""Kernel-level mask dimension support (2D, 3D, 4D)."""
|
||||
|
||||
import torch
|
||||
|
||||
from tests.extension.conftest import D, skip_no_kernel
|
||||
|
||||
|
||||
@skip_no_kernel
|
||||
def test_kernel_accepts_2d_mask():
|
||||
"""Kernel should accept 2D mask [batch, kv_len]."""
|
||||
from astrai.extension.attention_ops import attn_prefill
|
||||
|
||||
batch, q_len, n_heads, n_kv_heads = 1, 8, 4, 1
|
||||
kv_len = 8
|
||||
q = torch.randn(batch, q_len, n_heads, D, device="cuda", dtype=torch.bfloat16)
|
||||
k = torch.randn(batch, kv_len, n_kv_heads, D, device="cuda", dtype=torch.bfloat16)
|
||||
v = torch.randn(batch, kv_len, n_kv_heads, D, device="cuda", dtype=torch.bfloat16)
|
||||
mask = torch.ones(batch, kv_len, dtype=torch.bool, device="cuda")
|
||||
mask[:, 4:] = False
|
||||
|
||||
out = attn_prefill(q, k, v, mask=mask, is_causal=False)
|
||||
assert out.shape == (batch, q_len, n_heads, D)
|
||||
|
||||
|
||||
@skip_no_kernel
|
||||
def test_kernel_accepts_3d_mask():
|
||||
"""Kernel should accept 3D mask [batch, q_len, kv_len]."""
|
||||
from astrai.extension.attention_ops import attn_prefill
|
||||
|
||||
batch, q_len, n_heads, n_kv_heads = 1, 8, 4, 1
|
||||
kv_len = 8
|
||||
q = torch.randn(batch, q_len, n_heads, D, device="cuda", dtype=torch.bfloat16)
|
||||
k = torch.randn(batch, kv_len, n_kv_heads, D, device="cuda", dtype=torch.bfloat16)
|
||||
v = torch.randn(batch, kv_len, n_kv_heads, D, device="cuda", dtype=torch.bfloat16)
|
||||
mask = torch.ones(batch, q_len, kv_len, dtype=torch.bool, device="cuda")
|
||||
|
||||
out = attn_prefill(q, k, v, mask=mask, is_causal=False)
|
||||
assert out.shape == (batch, q_len, n_heads, D)
|
||||
|
||||
|
||||
@skip_no_kernel
|
||||
def test_kernel_accepts_4d_mask():
|
||||
"""Kernel should accept 4D mask [batch, n_heads, q_len, kv_len]."""
|
||||
from astrai.extension.attention_ops import attn_prefill
|
||||
|
||||
batch, q_len, n_heads, n_kv_heads = 1, 8, 4, 1
|
||||
kv_len = 8
|
||||
q = torch.randn(batch, q_len, n_heads, D, device="cuda", dtype=torch.bfloat16)
|
||||
k = torch.randn(batch, kv_len, n_kv_heads, D, device="cuda", dtype=torch.bfloat16)
|
||||
v = torch.randn(batch, kv_len, n_kv_heads, D, device="cuda", dtype=torch.bfloat16)
|
||||
mask = torch.ones(batch, 1, q_len, kv_len, dtype=torch.bool, device="cuda")
|
||||
mask[:, :, :, 4:] = False
|
||||
|
||||
out = attn_prefill(q, k, v, mask=mask, is_causal=False)
|
||||
assert out.shape == (batch, q_len, n_heads, D)
|
||||
|
||||
|
||||
@skip_no_kernel
|
||||
def test_4d_mask_matches_no_mask_when_all_true():
|
||||
"""A 4D all-True mask should produce the same output as no mask."""
|
||||
from astrai.extension.attention_ops import attn_prefill
|
||||
|
||||
batch, q_len, n_heads, n_kv_heads = 1, 8, 4, 1
|
||||
kv_len = 8
|
||||
q = torch.randn(batch, q_len, n_heads, D, device="cuda", dtype=torch.bfloat16)
|
||||
k = torch.randn(batch, kv_len, n_kv_heads, D, device="cuda", dtype=torch.bfloat16)
|
||||
v = torch.randn(batch, kv_len, n_kv_heads, D, device="cuda", dtype=torch.bfloat16)
|
||||
|
||||
out_no_mask = attn_prefill(q, k, v, mask=None, is_causal=False)
|
||||
mask = torch.ones(batch, 1, q_len, kv_len, dtype=torch.bool, device="cuda")
|
||||
out_with_mask = attn_prefill(q, k, v, mask=mask, is_causal=False)
|
||||
|
||||
diff = (out_no_mask.float() - out_with_mask.float()).abs().max().item()
|
||||
assert diff == 0.0, f"4D all-True mask diff: {diff}"
|
||||
@@ -0,0 +1,209 @@
|
||||
"""Shared test helpers for the AstrAI test suite."""
|
||||
|
||||
import json
|
||||
import os
|
||||
|
||||
import torch
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
|
||||
TINY_CONFIG = dict(
|
||||
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,
|
||||
)
|
||||
|
||||
CHAT_TEMPLATE = (
|
||||
"{% for message in messages %}"
|
||||
"{% if message['role'] == 'system' %}SYSTEM: {{ message['content'] }}\n{% endif %}"
|
||||
"{% if message['role'] == 'user' %}USER: {{ message['content'] }}\n{% endif %}"
|
||||
"{% if message['role'] == 'assistant' %}ASSISTANT: {{ message['content'] }}\n{% endif %}"
|
||||
"{% endfor %}"
|
||||
"{% if add_generation_prompt %}ASSISTANT: {% endif %}"
|
||||
)
|
||||
|
||||
|
||||
def make_tiny_config(**overrides):
|
||||
"""Create a tiny ``AutoRegressiveLMConfig`` for tests.
|
||||
|
||||
All keyword arguments override ``TINY_CONFIG`` defaults.
|
||||
"""
|
||||
return AutoRegressiveLMConfig(**{**TINY_CONFIG, **overrides})
|
||||
|
||||
|
||||
def make_rollout_config(vocab_size=200, max_position_embeddings=64, **kwargs):
|
||||
"""Create a tiny config sized for rollout / strategy tests."""
|
||||
return make_tiny_config(
|
||||
vocab_size=vocab_size,
|
||||
hidden_size=16,
|
||||
intermediate_size=32,
|
||||
max_position_embeddings=max_position_embeddings,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
def make_model(device, **cfg_overrides):
|
||||
"""Create a tiny ``AutoRegressiveLM`` on *device* and return ``(model, config)``."""
|
||||
cfg = make_rollout_config(**cfg_overrides)
|
||||
model = AutoRegressiveLM(cfg).to(device=device)
|
||||
model.eval()
|
||||
return model, cfg
|
||||
|
||||
|
||||
def make_frozen(model, device):
|
||||
"""Create a frozen, eval-mode copy of *model* with identical weights."""
|
||||
cfg = make_rollout_config()
|
||||
copy = AutoRegressiveLM(cfg).to(device=device)
|
||||
copy.load_state_dict(model.state_dict())
|
||||
copy.requires_grad_(False)
|
||||
copy.eval()
|
||||
return copy
|
||||
|
||||
|
||||
class RandomTokenDataset(Dataset):
|
||||
"""Random token dataset combining all test dataset variants.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
length : int or None
|
||||
Fixed length, or ``None`` for a random length in [100, 200).
|
||||
max_length : int
|
||||
Sequence length per sample.
|
||||
vocab_size : int
|
||||
Upper bound for random token ids.
|
||||
with_loss_mask : bool
|
||||
Include a ``loss_mask`` key in each sample.
|
||||
stop_after : int or None
|
||||
Raise ``RuntimeError`` after this many samples (for early-stopping tests).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
length=100,
|
||||
max_length=64,
|
||||
vocab_size=1000,
|
||||
*,
|
||||
with_loss_mask=False,
|
||||
stop_after=None,
|
||||
):
|
||||
self.length = (
|
||||
length if length is not None else int(torch.randint(100, 200, (1,)).item())
|
||||
)
|
||||
self.max_length = max_length
|
||||
self.vocab_size = vocab_size
|
||||
self.with_loss_mask = with_loss_mask
|
||||
self.stop_after = stop_after
|
||||
self._count = 0
|
||||
|
||||
def __len__(self):
|
||||
return self.length
|
||||
|
||||
def __getitem__(self, idx):
|
||||
if self.stop_after is not None:
|
||||
self._count += 1
|
||||
if self._count == self.stop_after:
|
||||
raise RuntimeError("Simulated early stopping")
|
||||
|
||||
item = {
|
||||
"input_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
|
||||
"target_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
|
||||
}
|
||||
if self.with_loss_mask:
|
||||
item["loss_mask"] = torch.randint(0, 1, (self.max_length,))
|
||||
return item
|
||||
|
||||
|
||||
class FakeTokenizer:
|
||||
"""Minimal stub tokenizer with optional chat-template support."""
|
||||
|
||||
stop_ids = [2]
|
||||
|
||||
def __init__(self, *, with_chat_template=False):
|
||||
if with_chat_template:
|
||||
from astrai.tokenize.chat_template import ChatTemplate
|
||||
|
||||
self._chat_template = ChatTemplate.from_string(CHAT_TEMPLATE)
|
||||
else:
|
||||
self._chat_template = None
|
||||
|
||||
def encode(self, texts, **_):
|
||||
if isinstance(texts, str):
|
||||
texts = [texts]
|
||||
return [[b for b in t.encode("utf-8")] for t in texts]
|
||||
|
||||
def decode(self, ids, skip_special_tokens=True):
|
||||
if isinstance(ids, list):
|
||||
return bytes(b for b in ids if b > 2 or not skip_special_tokens).decode(
|
||||
"utf-8", errors="ignore"
|
||||
)
|
||||
return str(ids)
|
||||
|
||||
def apply_chat_template(
|
||||
self, messages, tokenize=True, add_generation_prompt=True, **_
|
||||
):
|
||||
if self._chat_template is None:
|
||||
raise RuntimeError("Chat template not configured")
|
||||
rendered = self._chat_template.render(
|
||||
messages=messages, add_generation_prompt=add_generation_prompt
|
||||
)
|
||||
if tokenize:
|
||||
return (
|
||||
self.encode(rendered)[0]
|
||||
if isinstance(rendered, str)
|
||||
else [self.encode(t)[0] for t in rendered]
|
||||
)
|
||||
return rendered
|
||||
|
||||
|
||||
class FakeExecutor:
|
||||
"""Executor stub tracking ``sync_gradients`` and providing ``unwrap_model``."""
|
||||
|
||||
use_distributed = False
|
||||
|
||||
def __init__(self, sync_gradients=True):
|
||||
self._sync_gradients = sync_gradients
|
||||
|
||||
@property
|
||||
def sync_gradients(self):
|
||||
return self._sync_gradients
|
||||
|
||||
def unwrap_model(self, model):
|
||||
return model.state_dict()
|
||||
|
||||
|
||||
def find_checkpoint_meta(ckpt_dir):
|
||||
"""Walk *ckpt_dir* and return the path to the first ``meta.json`` found."""
|
||||
for root, _dirs, files in os.walk(ckpt_dir):
|
||||
if "meta.json" in files:
|
||||
return os.path.join(root, "meta.json")
|
||||
return None
|
||||
|
||||
|
||||
def load_checkpoint_meta(ckpt_dir):
|
||||
"""Find and load the first checkpoint ``meta.json`` under *ckpt_dir*."""
|
||||
meta_path = find_checkpoint_meta(ckpt_dir)
|
||||
assert meta_path is not None, f"No checkpoint meta.json found in {ckpt_dir}"
|
||||
with open(meta_path) as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def load_shard_meta(out_dir):
|
||||
"""Load ``meta.json`` from the default shard output directory."""
|
||||
meta_path = os.path.join(out_dir, "__default__", "shard_0000", "meta.json")
|
||||
assert os.path.exists(meta_path), f"Shard meta.json not found at {meta_path}"
|
||||
with open(meta_path) as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def assert_state_dicts_equal(a, b):
|
||||
"""Assert two state dicts have identical keys and equal tensor values."""
|
||||
assert set(a.keys()) == set(b.keys()), f"Key mismatch: {set(a) ^ set(b)}"
|
||||
for key in a:
|
||||
assert torch.equal(a[key], b[key]), f"Tensor mismatch at key: {key}"
|
||||
+270
-189
@@ -4,17 +4,14 @@ import torch
|
||||
|
||||
from astrai.inference import (
|
||||
Allocator,
|
||||
PageCache,
|
||||
KVStorage,
|
||||
PagePool,
|
||||
PrefixCache,
|
||||
Storage,
|
||||
TaskTable,
|
||||
ReqToTokenPool,
|
||||
page_hash,
|
||||
)
|
||||
|
||||
|
||||
def make_pool(n_pages: int, page_size: int) -> PagePool:
|
||||
return PagePool(Allocator(n_pages), PrefixCache(page_size))
|
||||
# ---- page_hash ----
|
||||
|
||||
|
||||
def test_page_hash_full_page():
|
||||
@@ -29,251 +26,335 @@ def test_page_hash_different_page_differs():
|
||||
assert page_hash(token_ids, 0, 64) != page_hash(token_ids, 1, 64)
|
||||
|
||||
|
||||
def test_page_pool_alloc_free_cycle():
|
||||
pool = make_pool(4, 64)
|
||||
a = pool.alloc()
|
||||
b = pool.alloc()
|
||||
# ---- Allocator ----
|
||||
|
||||
|
||||
def test_allocator_alloc_free_cycle():
|
||||
alloc = Allocator(4)
|
||||
a = alloc.alloc()
|
||||
b = alloc.alloc()
|
||||
assert a != b
|
||||
pool.free(a)
|
||||
pool.free(b)
|
||||
c = pool.alloc()
|
||||
alloc.free(a)
|
||||
alloc.free(b)
|
||||
c = alloc.alloc()
|
||||
assert c in (a, b)
|
||||
|
||||
|
||||
def test_page_pool_alloc_when_full():
|
||||
pool = make_pool(2, 64)
|
||||
pool.alloc()
|
||||
pool.alloc()
|
||||
assert pool.alloc() == -1
|
||||
def test_allocator_alloc_when_full():
|
||||
alloc = Allocator(2)
|
||||
alloc.alloc()
|
||||
alloc.alloc()
|
||||
assert alloc.alloc() == -1
|
||||
|
||||
|
||||
def test_page_pool_lru_eviction():
|
||||
pool = make_pool(2, 64)
|
||||
p0 = pool.alloc()
|
||||
p1 = pool.alloc()
|
||||
pool.record(p0, list(range(64)), 0)
|
||||
pool.record(p1, list(range(64, 128)), 0)
|
||||
pool.free(p0)
|
||||
pool.free(p1)
|
||||
pool.alloc()
|
||||
assert p0 in pool._alloc._lru or p1 in pool._alloc._lru
|
||||
def test_allocator_lru_eviction():
|
||||
alloc = Allocator(2)
|
||||
p0 = alloc.alloc()
|
||||
p1 = alloc.alloc()
|
||||
alloc.free(p0, keep_cached=True)
|
||||
alloc.free(p1, keep_cached=True)
|
||||
alloc.alloc()
|
||||
assert p0 in alloc._lru or p1 in alloc._lru
|
||||
|
||||
|
||||
def test_page_pool_inc_ref_and_free():
|
||||
pool = make_pool(2, 64)
|
||||
p = pool.alloc()
|
||||
pool.inc_ref(p)
|
||||
assert pool._alloc._refs[p] == 2
|
||||
pool.free(p)
|
||||
assert pool._alloc._refs[p] == 1
|
||||
pool.free(p)
|
||||
assert pool._alloc._refs[p] == 0
|
||||
def test_allocator_inc_ref_and_free():
|
||||
alloc = Allocator(2)
|
||||
p = alloc.alloc()
|
||||
alloc.inc_ref(p)
|
||||
assert alloc._refs[p] == 2
|
||||
alloc.free(p)
|
||||
assert alloc._refs[p] == 1
|
||||
alloc.free(p)
|
||||
assert alloc._refs[p] == 0
|
||||
|
||||
|
||||
def test_page_pool_keep_cached_realloc():
|
||||
"""Free mask has priority over LRU; cached page returned only when no free pages."""
|
||||
pool = make_pool(3, 64)
|
||||
p0 = pool.alloc()
|
||||
p1 = pool.alloc()
|
||||
p2 = pool.alloc()
|
||||
for p in (p0, p1, p2):
|
||||
pool.record(p, [p] * 64, 0)
|
||||
pool.free(p0)
|
||||
pool.free(p1)
|
||||
pool.free(p2)
|
||||
assert pool.alloc() == p0
|
||||
# ---- PrefixCache ----
|
||||
|
||||
|
||||
def test_prefix_cache_lookup_returns_hits():
|
||||
token_ids = list(range(256))
|
||||
pool = make_pool(16, 64)
|
||||
pages = [pool.alloc() for _ in range(4)]
|
||||
prefix = PrefixCache(64)
|
||||
pages = [0, 1, 2, 3]
|
||||
for i, p in enumerate(pages):
|
||||
pool.record(p, token_ids, i)
|
||||
pool.free(p)
|
||||
hits = pool.lookup(token_ids)
|
||||
prefix.record(p, token_ids, i)
|
||||
hits = prefix.lookup(token_ids)
|
||||
assert hits == pages
|
||||
|
||||
|
||||
def test_prefix_cache_lookup_stops_at_first_miss():
|
||||
token_ids = list(range(256))
|
||||
pool = make_pool(16, 64)
|
||||
p0 = pool.alloc()
|
||||
pool.record(p0, token_ids, 0)
|
||||
pool.free(p0)
|
||||
p1 = pool.alloc()
|
||||
pool.record(p1, [99] * 64, 1)
|
||||
pool.free(p1)
|
||||
hits = pool.lookup(token_ids)
|
||||
prefix = PrefixCache(64)
|
||||
prefix.record(0, token_ids, 0)
|
||||
prefix.record(1, [99] * 64, 1)
|
||||
hits = prefix.lookup(token_ids)
|
||||
assert len(hits) == 1
|
||||
assert hits[0] == p0
|
||||
assert hits[0] == 0
|
||||
|
||||
|
||||
def test_prefix_cache_ignores_partial_last_page():
|
||||
token_ids = list(range(100))
|
||||
pool = make_pool(16, 64)
|
||||
p = pool.alloc()
|
||||
pool.record(p, token_ids, 0)
|
||||
pool.free(p)
|
||||
hits = pool.lookup(token_ids)
|
||||
prefix = PrefixCache(64)
|
||||
prefix.record(0, token_ids, 0)
|
||||
hits = prefix.lookup(token_ids)
|
||||
assert len(hits) == 1
|
||||
|
||||
|
||||
def test_prefix_cache_on_evict_clears_mappings():
|
||||
pool = make_pool(4, 64)
|
||||
p = pool.alloc()
|
||||
pool.record(p, list(range(64)), 0)
|
||||
pool.free(p)
|
||||
assert p in pool._prefix._page_to_hash
|
||||
pool._prefix.evict(p)
|
||||
assert p not in pool._prefix._page_to_hash
|
||||
prefix = PrefixCache(64)
|
||||
prefix.record(0, list(range(64)), 0)
|
||||
assert 0 in prefix._page_to_hash
|
||||
prefix.evict(0)
|
||||
assert 0 not in prefix._page_to_hash
|
||||
|
||||
|
||||
def test_prefix_cache_has_page():
|
||||
pool = make_pool(4, 64)
|
||||
p = pool.alloc()
|
||||
assert p not in pool._prefix._page_to_hash
|
||||
pool.record(p, list(range(64)), 0)
|
||||
pool.free(p)
|
||||
assert p in pool._prefix._page_to_hash
|
||||
prefix = PrefixCache(64)
|
||||
assert not prefix.has_page(0)
|
||||
prefix.record(0, list(range(64)), 0)
|
||||
assert prefix.has_page(0)
|
||||
|
||||
|
||||
def test_task_table_set_get():
|
||||
table = TaskTable(page_size=64)
|
||||
table.set("task1", [0, 1, 2], 128)
|
||||
assert table.get("task1") == [0, 1, 2]
|
||||
assert table.get_cached("task1") == 128
|
||||
# ---- ReqToTokenPool ----
|
||||
|
||||
|
||||
def test_task_table_get_missing():
|
||||
table = TaskTable(page_size=64)
|
||||
assert table.get("nonexistent") == []
|
||||
assert table.get_cached("nonexistent") == 0
|
||||
def test_req_to_token_pool_alloc_free():
|
||||
pool = ReqToTokenPool(4, 128, torch.device("cpu"))
|
||||
slots = pool.alloc(2)
|
||||
assert len(slots) == 2
|
||||
assert len(pool.free_slots) == 2
|
||||
pool.free(slots)
|
||||
assert len(pool.free_slots) == 4
|
||||
|
||||
|
||||
def test_task_table_pop():
|
||||
table = TaskTable(page_size=64)
|
||||
table.set("task1", [0, 1], 64)
|
||||
pages, cached = table.pop("task1")
|
||||
assert pages == [0, 1]
|
||||
assert cached == 64
|
||||
assert table.get("task1") == []
|
||||
def test_req_to_token_pool_alloc_when_full():
|
||||
pool = ReqToTokenPool(2, 128, torch.device("cpu"))
|
||||
pool.alloc(2)
|
||||
assert pool.alloc(1) is None
|
||||
|
||||
|
||||
def test_kv_cache_task_extend_allocates():
|
||||
cache = PageCache(
|
||||
n_layers=1,
|
||||
n_pages=8,
|
||||
page_size=64,
|
||||
n_kv_heads=2,
|
||||
head_dim=8,
|
||||
device=torch.device("cpu"),
|
||||
dtype=torch.float32,
|
||||
)
|
||||
cache._table.set("task1", [], 0)
|
||||
ok = cache.task_extend("task1", 200)
|
||||
assert ok
|
||||
assert len(cache._table.get("task1")) == 4
|
||||
def test_req_to_token_pool_write():
|
||||
pool = ReqToTokenPool(4, 128, torch.device("cpu"))
|
||||
slots = pool.alloc(1)
|
||||
pool.write((slots[0], slice(0, 3)), torch.tensor([10, 20, 30]))
|
||||
assert pool.req_to_token[slots[0], 0].item() == 10
|
||||
assert pool.req_to_token[slots[0], 2].item() == 30
|
||||
|
||||
|
||||
def test_kv_cache_task_extend_fails_when_pool_full():
|
||||
cache = PageCache(
|
||||
n_layers=1,
|
||||
n_pages=2,
|
||||
page_size=64,
|
||||
n_kv_heads=2,
|
||||
head_dim=8,
|
||||
device=torch.device("cpu"),
|
||||
dtype=torch.float32,
|
||||
)
|
||||
cache._table.set("task1", [0, 1], 0)
|
||||
ok = cache.task_extend("task1", 300)
|
||||
assert not ok
|
||||
# ---- KVStorage ----
|
||||
|
||||
|
||||
def test_task_table_table_tensor():
|
||||
table = TaskTable(page_size=64)
|
||||
table.set("a", [0, 1], 0)
|
||||
table.set("b", [2, 3, 4], 0)
|
||||
t = table.table_tensor(["a", "b"], torch.device("cpu"))
|
||||
assert t.shape == (2, 3)
|
||||
assert t[0].tolist() == [0, 1, -1]
|
||||
assert t[1].tolist() == [2, 3, 4]
|
||||
|
||||
|
||||
def test_task_table_table_tensor_empty_input():
|
||||
table = TaskTable(page_size=64)
|
||||
t = table.table_tensor([], torch.device("cpu"))
|
||||
assert t.numel() == 0
|
||||
|
||||
|
||||
def test_storage_write_gather_single_page():
|
||||
storage = Storage(
|
||||
def test_kv_storage_set_and_get():
|
||||
storage = KVStorage(
|
||||
size=16,
|
||||
n_layers=2,
|
||||
n_pages=8,
|
||||
page_size=4,
|
||||
n_kv_heads=2,
|
||||
n_kv_heads=4,
|
||||
head_dim=8,
|
||||
device=torch.device("cpu"),
|
||||
dtype=torch.float32,
|
||||
)
|
||||
page_table = torch.tensor([[0]], dtype=torch.long)
|
||||
k = torch.randn(1, 2, 2, 8)
|
||||
v = torch.randn(1, 2, 2, 8)
|
||||
|
||||
storage.write(0, page_table, 0, k, v)
|
||||
gk, gv = storage.gather(0, page_table, 2)
|
||||
assert torch.allclose(gk, k)
|
||||
loc = torch.tensor([[0, 1]], dtype=torch.long)
|
||||
k = torch.randn(1, 2, 4, 8)
|
||||
v = torch.randn(1, 2, 4, 8)
|
||||
storage.set_kv_buffer(0, loc, k, v)
|
||||
assert torch.allclose(storage.get_key_buffer(0)[loc], k)
|
||||
assert torch.allclose(storage.get_value_buffer(0)[loc], v)
|
||||
|
||||
|
||||
def test_storage_write_cross_page():
|
||||
storage = Storage(
|
||||
def test_kv_storage_buffer_shape():
|
||||
storage = KVStorage(
|
||||
size=32,
|
||||
n_layers=3,
|
||||
n_kv_heads=8,
|
||||
head_dim=16,
|
||||
device=torch.device("cpu"),
|
||||
dtype=torch.float32,
|
||||
)
|
||||
assert storage.k_buffer.shape == (3, 32, 8, 16)
|
||||
assert storage.v_buffer.shape == (3, 32, 8, 16)
|
||||
|
||||
|
||||
# ---- PagePool (contiguous mode) ----
|
||||
|
||||
|
||||
def _make_contiguous_pool(**kwargs):
|
||||
defaults = dict(
|
||||
n_layers=2,
|
||||
n_kv_heads=4,
|
||||
head_dim=8,
|
||||
max_batch_size=4,
|
||||
max_seq_len=64,
|
||||
device=torch.device("cpu"),
|
||||
dtype=torch.float32,
|
||||
)
|
||||
defaults.update(kwargs)
|
||||
return PagePool(**defaults)
|
||||
|
||||
|
||||
def test_page_pool_contiguous_task_alloc_free():
|
||||
pool = _make_contiguous_pool()
|
||||
assert pool.task_alloc("t1", [1, 2, 3])
|
||||
assert "t1" in pool._task_req
|
||||
pool.task_free("t1")
|
||||
assert "t1" not in pool._task_req
|
||||
|
||||
|
||||
def test_page_pool_contiguous_task_extend():
|
||||
pool = _make_contiguous_pool()
|
||||
pool.task_alloc("t1", [1, 2, 3])
|
||||
assert pool.task_extend("t1", 3)
|
||||
assert pool.task_extend("t1", 63)
|
||||
assert not pool.task_extend("t1", 64)
|
||||
|
||||
|
||||
def test_page_pool_contiguous_task_cached():
|
||||
pool = _make_contiguous_pool()
|
||||
pool.task_alloc("t1", [1, 2, 3])
|
||||
assert pool.task_cached("t1") == 0
|
||||
|
||||
|
||||
def test_page_pool_contiguous_bind_tasks_prefill():
|
||||
pool = _make_contiguous_pool()
|
||||
pool.task_alloc("t1", list(range(10)))
|
||||
pool.task_alloc("t2", list(range(10)))
|
||||
kv = pool.bind_tasks(["t1", "t2"], [10, 10], torch.device("cpu"), start_pos=0)
|
||||
assert kv.out_cache_loc.shape == (2, 10)
|
||||
assert kv.seq_lens.tolist() == [10, 10]
|
||||
assert kv.req_pool_indices.shape == (2,)
|
||||
|
||||
|
||||
def test_page_pool_contiguous_bind_tasks_decode():
|
||||
pool = _make_contiguous_pool()
|
||||
pool.task_alloc("t1", list(range(10)))
|
||||
pool.task_alloc("t2", list(range(8)))
|
||||
kv = pool.bind_tasks(["t1", "t2"], [11, 9], torch.device("cpu"))
|
||||
assert kv.out_cache_loc.shape == (2, 1)
|
||||
assert kv.seq_lens.tolist() == [11, 9]
|
||||
|
||||
|
||||
def test_page_pool_contiguous_bind_roundtrip():
|
||||
"""Write KV via bind_tasks, then gather via req_to_token indexing."""
|
||||
pool = _make_contiguous_pool(n_layers=1, n_kv_heads=2, head_dim=4)
|
||||
pool.task_alloc("t1", list(range(4)))
|
||||
|
||||
kv = pool.bind_tasks(["t1"], [4], torch.device("cpu"), start_pos=0)
|
||||
k = torch.randn(1, 4, 2, 4)
|
||||
v = torch.randn(1, 4, 2, 4)
|
||||
kv.k_buffer[0, kv.out_cache_loc] = k
|
||||
kv.v_buffer[0, kv.out_cache_loc] = v
|
||||
|
||||
indices = kv.req_to_token[kv.req_pool_indices, :4]
|
||||
gathered_k = kv.k_buffer[0, indices]
|
||||
gathered_v = kv.v_buffer[0, indices]
|
||||
assert torch.allclose(gathered_k, k)
|
||||
assert torch.allclose(gathered_v, v)
|
||||
|
||||
|
||||
# ---- PagePool (paged mode, page_size=1) ----
|
||||
|
||||
|
||||
def _make_paged_pool(**kwargs):
|
||||
defaults = dict(
|
||||
n_layers=1,
|
||||
n_pages=8,
|
||||
page_size=4,
|
||||
n_kv_heads=2,
|
||||
head_dim=8,
|
||||
head_dim=4,
|
||||
max_batch_size=4,
|
||||
max_seq_len=64,
|
||||
device=torch.device("cpu"),
|
||||
dtype=torch.float32,
|
||||
page_size=1,
|
||||
n_tokens=128,
|
||||
)
|
||||
page_table = torch.tensor([[0, 1]], dtype=torch.long)
|
||||
k = torch.randn(1, 8, 2, 8)
|
||||
v = torch.randn(1, 8, 2, 8)
|
||||
|
||||
storage.write(0, page_table, 0, k, v)
|
||||
gk, gv = storage.gather(0, page_table, 8)
|
||||
assert torch.allclose(gk, k)
|
||||
defaults.update(kwargs)
|
||||
return PagePool(**defaults)
|
||||
|
||||
|
||||
def test_storage_gather_truncates_to_total_len():
|
||||
storage = Storage(
|
||||
def test_page_pool_paged_task_alloc():
|
||||
pool = _make_paged_pool()
|
||||
assert pool.task_alloc("t1", list(range(10)))
|
||||
req_idx = pool._task_req["t1"]
|
||||
slots = pool._task_slots["t1"]
|
||||
assert len(slots) == 10
|
||||
assert pool._req_pool.req_to_token[req_idx, 0].item() == slots[0]
|
||||
|
||||
|
||||
def test_page_pool_paged_task_extend():
|
||||
pool = _make_paged_pool()
|
||||
pool.task_alloc("t1", list(range(4)))
|
||||
assert pool.task_extend("t1", 4)
|
||||
req_idx = pool._task_req["t1"]
|
||||
slot = pool._req_pool.req_to_token[req_idx, 4].item()
|
||||
assert slot >= 0
|
||||
|
||||
|
||||
def test_page_pool_paged_task_free_releases_slots():
|
||||
pool = _make_paged_pool(n_tokens=16)
|
||||
pool.task_alloc("t1", list(range(8)))
|
||||
pool.task_free("t1")
|
||||
assert "t1" not in pool._task_req
|
||||
assert len(pool._req_pool.free_slots) == 4
|
||||
|
||||
|
||||
def test_page_pool_paged_bind_roundtrip():
|
||||
pool = _make_paged_pool(n_layers=1, n_kv_heads=2, head_dim=4)
|
||||
pool.task_alloc("t1", list(range(4)))
|
||||
|
||||
kv = pool.bind_tasks(["t1"], [4], torch.device("cpu"), start_pos=0)
|
||||
k = torch.randn(1, 4, 2, 4)
|
||||
v = torch.randn(1, 4, 2, 4)
|
||||
kv.k_buffer[0, kv.out_cache_loc] = k
|
||||
kv.v_buffer[0, kv.out_cache_loc] = v
|
||||
|
||||
indices = kv.req_to_token[kv.req_pool_indices, :4]
|
||||
gathered_k = kv.k_buffer[0, indices]
|
||||
assert torch.allclose(gathered_k, k)
|
||||
|
||||
|
||||
# ---- PagePool (paged mode, page_size>1) ----
|
||||
|
||||
|
||||
def _make_paged_pool_ps64(**kwargs):
|
||||
defaults = dict(
|
||||
n_layers=1,
|
||||
n_pages=8,
|
||||
page_size=4,
|
||||
n_kv_heads=2,
|
||||
head_dim=8,
|
||||
head_dim=4,
|
||||
max_batch_size=4,
|
||||
max_seq_len=256,
|
||||
device=torch.device("cpu"),
|
||||
dtype=torch.float32,
|
||||
page_size=64,
|
||||
n_tokens=512,
|
||||
)
|
||||
page_table = torch.tensor([[0, 1]], dtype=torch.long)
|
||||
k = torch.randn(1, 6, 2, 8)
|
||||
v = torch.randn(1, 6, 2, 8)
|
||||
storage.write(0, page_table, 0, k, v)
|
||||
|
||||
gk, gv = storage.gather(0, page_table, 5)
|
||||
assert gk.shape == (1, 5, 2, 8)
|
||||
defaults.update(kwargs)
|
||||
return PagePool(**defaults)
|
||||
|
||||
|
||||
def test_storage_gather_clamps_negative_padding():
|
||||
storage = Storage(
|
||||
n_layers=1,
|
||||
n_pages=8,
|
||||
page_size=4,
|
||||
n_kv_heads=2,
|
||||
head_dim=8,
|
||||
device=torch.device("cpu"),
|
||||
dtype=torch.float32,
|
||||
)
|
||||
page_table = torch.tensor([[0, -1]], dtype=torch.long)
|
||||
gk, gv = storage.gather(0, page_table, 4)
|
||||
assert gk.shape == (1, 4, 2, 8)
|
||||
def test_page_pool_paged_ps64_task_alloc():
|
||||
pool = _make_paged_pool_ps64()
|
||||
prompt = list(range(200))
|
||||
assert pool.task_alloc("t1", prompt)
|
||||
assert pool.task_cached("t1") == 0
|
||||
n_pages = (200 + 63) // 64
|
||||
assert len(pool._task_pages["t1"]) == n_pages
|
||||
|
||||
|
||||
def test_page_pool_paged_ps64_task_extend_crosses_page():
|
||||
pool = _make_paged_pool_ps64()
|
||||
pool.task_alloc("t1", list(range(64)))
|
||||
assert pool.task_extend("t1", 64)
|
||||
assert len(pool._task_pages["t1"]) >= 2
|
||||
|
||||
|
||||
def test_page_pool_paged_ps64_bind_roundtrip():
|
||||
pool = _make_paged_pool_ps64(n_layers=1, n_kv_heads=2, head_dim=4)
|
||||
prompt = list(range(128))
|
||||
pool.task_alloc("t1", prompt)
|
||||
|
||||
kv = pool.bind_tasks(["t1"], [128], torch.device("cpu"), start_pos=0)
|
||||
k = torch.randn(1, 128, 2, 4)
|
||||
v = torch.randn(1, 128, 2, 4)
|
||||
kv.k_buffer[0, kv.out_cache_loc] = k
|
||||
kv.v_buffer[0, kv.out_cache_loc] = v
|
||||
|
||||
indices = kv.req_to_token[kv.req_pool_indices, :128]
|
||||
gathered_k = kv.k_buffer[0, indices]
|
||||
assert torch.allclose(gathered_k, k)
|
||||
|
||||
@@ -7,6 +7,8 @@ import pytest
|
||||
import torch
|
||||
|
||||
from astrai.inference import InferenceScheduler
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
from tests.helpers import FakeTokenizer, make_rollout_config
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -195,46 +197,19 @@ def test_prefill_skips_fully_cached_tasks(mock_model_and_tokenizer):
|
||||
|
||||
def _make_real_scheduler(device):
|
||||
"""Build a scheduler backed by a tiny real model for run_batch tests."""
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
|
||||
class _Tok:
|
||||
stop_ids = [2]
|
||||
|
||||
def encode(self, texts, **_):
|
||||
if isinstance(texts, str):
|
||||
texts = [texts]
|
||||
return [[b for b in t.encode("utf-8")] for t in texts]
|
||||
|
||||
def decode(self, ids, skip_special_tokens=True):
|
||||
return bytes(b for b in ids if b > 2 or not skip_special_tokens).decode(
|
||||
"utf-8", errors="ignore"
|
||||
)
|
||||
|
||||
cfg = AutoRegressiveLMConfig(
|
||||
vocab_size=200,
|
||||
hidden_size=16,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=1,
|
||||
intermediate_size=32,
|
||||
max_position_embeddings=64,
|
||||
num_hidden_layers=2,
|
||||
rms_norm_eps=1e-5,
|
||||
)
|
||||
cfg = make_rollout_config(max_position_embeddings=64)
|
||||
model = AutoRegressiveLM(cfg).to(device=device).eval()
|
||||
tokenizer = _Tok()
|
||||
tokenizer = FakeTokenizer()
|
||||
scheduler = InferenceScheduler(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=8,
|
||||
max_seq_len=64,
|
||||
max_prompt_len=64,
|
||||
)
|
||||
return scheduler, tokenizer, model
|
||||
|
||||
|
||||
def test_run_batch_returns_token_sequences():
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
def test_run_batch_returns_token_sequences(device):
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
prompts = [[10, 20, 30], [5, 6, 7, 8]]
|
||||
@@ -248,9 +223,8 @@ def test_run_batch_returns_token_sequences():
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_run_batch_return_logprobs_aligned():
|
||||
def test_run_batch_return_logprobs_aligned(device):
|
||||
"""return_logprobs=True gives (token_ids, logprobs) tuples with equal len."""
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
prompts = [[10, 20, 30, 40]]
|
||||
@@ -265,8 +239,7 @@ def test_run_batch_return_logprobs_aligned():
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_run_batch_respects_max_tokens():
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
def test_run_batch_respects_max_tokens(device):
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
prompts = [[10, 20, 30]]
|
||||
@@ -276,9 +249,8 @@ def test_run_batch_respects_max_tokens():
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_run_batch_stop_id_terminates():
|
||||
def test_run_batch_stop_id_terminates(device):
|
||||
"""A token matching stop_ids terminates generation for that prompt."""
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
prompts = [[10, 20, 30]]
|
||||
@@ -291,9 +263,8 @@ def test_run_batch_stop_id_terminates():
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_run_batch_empty_prompts():
|
||||
def test_run_batch_empty_prompts(device):
|
||||
"""Empty prompt list yields empty result list."""
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
assert scheduler.run_batch([], max_tokens=4) == []
|
||||
@@ -301,9 +272,8 @@ def test_run_batch_empty_prompts():
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_run_batch_too_long_prompt_skipped():
|
||||
def test_run_batch_too_long_prompt_skipped(device):
|
||||
"""A prompt longer than max_seq_len yields an empty result slot."""
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
long = list(range(100)) # > max_seq_len=64
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from astrai.inference import STOP, Task, TaskManager, TaskStatus
|
||||
from astrai.inference import Task, TaskManager, TaskStatus
|
||||
|
||||
|
||||
def _make_mock_tokenizer():
|
||||
@@ -58,8 +58,8 @@ def test_task_manager_add_task_too_long_immediate_stop():
|
||||
|
||||
tm = TaskManager(tokenizer=t, max_seq_len=16)
|
||||
tm.add_task("long", stream_callback=lambda tok: cb_calls.append(tok))
|
||||
assert cb_calls[0] is STOP
|
||||
assert len(tm.waiting_queue) == 0
|
||||
assert len(cb_calls) == 0
|
||||
assert len(tm.waiting_queue) == 1
|
||||
|
||||
|
||||
def test_task_manager_remove_task():
|
||||
|
||||
@@ -7,36 +7,24 @@ import safetensors.torch as st
|
||||
import torch
|
||||
|
||||
from astrai.config.model_config import EncoderConfig
|
||||
from astrai.model.automodel import AutoModel
|
||||
from astrai.model.automodel import ModelFactory
|
||||
from astrai.model.encoder import EmbeddingEncoder
|
||||
|
||||
TINY_CONFIG = dict(
|
||||
vocab_size=128,
|
||||
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"
|
||||
from tests.helpers import TINY_CONFIG, assert_state_dicts_equal
|
||||
|
||||
|
||||
def _make_model(**kwargs):
|
||||
def _make_model(device, **kwargs):
|
||||
config = EncoderConfig(**{**TINY_CONFIG, **kwargs})
|
||||
return EmbeddingEncoder(config).to(device=_device)
|
||||
return EmbeddingEncoder(config).to(device=device)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("pooling_type", ["mean", "cls", "last"])
|
||||
def test_encoder_forward_pooling(pooling_type):
|
||||
model = _make_model(pooling_type=pooling_type)
|
||||
def test_encoder_forward_pooling(pooling_type, device):
|
||||
model = _make_model(device, pooling_type=pooling_type)
|
||||
model.eval()
|
||||
|
||||
batch_size, seq_len = 2, 8
|
||||
input_ids = torch.randint(
|
||||
0, TINY_CONFIG["vocab_size"], (batch_size, seq_len), device=_device
|
||||
0, TINY_CONFIG["vocab_size"], (batch_size, seq_len), device=device
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
@@ -46,15 +34,15 @@ def test_encoder_forward_pooling(pooling_type):
|
||||
assert not torch.isnan(output).any()
|
||||
|
||||
|
||||
def test_encoder_forward_with_padding():
|
||||
model = _make_model()
|
||||
def test_encoder_forward_with_padding(device):
|
||||
model = _make_model(device)
|
||||
model.eval()
|
||||
|
||||
batch_size, seq_len = 2, 8
|
||||
input_ids = torch.randint(
|
||||
0, TINY_CONFIG["vocab_size"], (batch_size, seq_len), device=_device
|
||||
0, TINY_CONFIG["vocab_size"], (batch_size, seq_len), device=device
|
||||
)
|
||||
input_mask = torch.ones(batch_size, seq_len, dtype=torch.bool, device=_device)
|
||||
input_mask = torch.ones(batch_size, seq_len, dtype=torch.bool, device=device)
|
||||
input_mask[:, 4:] = False
|
||||
|
||||
with torch.no_grad():
|
||||
@@ -64,13 +52,13 @@ def test_encoder_forward_with_padding():
|
||||
assert not torch.isnan(output).any()
|
||||
|
||||
|
||||
def test_encoder_normalize():
|
||||
model = _make_model(pooling_type="mean", normalize_embeddings=True)
|
||||
def test_encoder_normalize(device):
|
||||
model = _make_model(device, pooling_type="mean", normalize_embeddings=True)
|
||||
model.eval()
|
||||
|
||||
batch_size, seq_len = 2, 8
|
||||
input_ids = torch.randint(
|
||||
0, TINY_CONFIG["vocab_size"], (batch_size, seq_len), device=_device
|
||||
0, TINY_CONFIG["vocab_size"], (batch_size, seq_len), device=device
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
@@ -81,31 +69,29 @@ def test_encoder_normalize():
|
||||
|
||||
|
||||
def test_encoder_register():
|
||||
assert AutoModel.is_registered("embedding")
|
||||
cls = AutoModel.get_component_class("embedding")
|
||||
assert ModelFactory.is_registered("embedding")
|
||||
cls = ModelFactory.get_component_class("embedding")
|
||||
assert cls is EmbeddingEncoder
|
||||
|
||||
|
||||
def test_encoder_from_transformer_checkpoint():
|
||||
model = _make_model()
|
||||
def test_encoder_from_transformer_checkpoint(device):
|
||||
model = _make_model(device)
|
||||
state_dict = model.state_dict()
|
||||
state_dict["lm_head.weight"] = torch.randn(
|
||||
TINY_CONFIG["vocab_size"], TINY_CONFIG["hidden_size"], device=_device
|
||||
TINY_CONFIG["vocab_size"], TINY_CONFIG["hidden_size"], device=device
|
||||
)
|
||||
|
||||
new_model = _make_model()
|
||||
new_model = _make_model(device)
|
||||
new_model.load_state_dict(state_dict, strict=True)
|
||||
|
||||
for key in model.state_dict():
|
||||
assert torch.equal(new_model.state_dict()[key], model.state_dict()[key])
|
||||
assert_state_dicts_equal(new_model.state_dict(), model.state_dict())
|
||||
|
||||
|
||||
def test_encoder_save_load():
|
||||
test_dir = tempfile.mkdtemp(prefix="encoder_test_")
|
||||
config_path = os.path.join(test_dir, "config.json")
|
||||
weights_path = os.path.join(test_dir, "model.safetensors")
|
||||
def test_encoder_save_load(device):
|
||||
with tempfile.TemporaryDirectory(prefix="encoder_test_") as test_dir:
|
||||
config_path = os.path.join(test_dir, "config.json")
|
||||
weights_path = os.path.join(test_dir, "model.safetensors")
|
||||
|
||||
try:
|
||||
config_data = {**TINY_CONFIG, "pooling_type": "mean"}
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_data, f)
|
||||
@@ -117,10 +103,4 @@ def test_encoder_save_load():
|
||||
loaded = EmbeddingEncoder(config)
|
||||
loaded.load_state_dict(st.load_file(weights_path))
|
||||
|
||||
for key in original.state_dict():
|
||||
assert torch.equal(original.state_dict()[key], loaded.state_dict()[key])
|
||||
finally:
|
||||
if os.path.exists(test_dir):
|
||||
for f in os.listdir(test_dir):
|
||||
os.remove(os.path.join(test_dir, f))
|
||||
os.rmdir(test_dir)
|
||||
assert_state_dicts_equal(original.state_dict(), loaded.state_dict())
|
||||
|
||||
@@ -1,20 +1,8 @@
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
|
||||
TINY_CONFIG = dict(
|
||||
vocab_size=128,
|
||||
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,
|
||||
)
|
||||
|
||||
from tests.helpers import TINY_CONFIG
|
||||
|
||||
CONFIGS = [
|
||||
pytest.param(
|
||||
@@ -70,9 +58,10 @@ CONFIGS = [
|
||||
|
||||
|
||||
@pytest.mark.parametrize("config_kwargs", CONFIGS)
|
||||
def test_model_forward(config_kwargs):
|
||||
def test_model_forward(config_kwargs, device):
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
|
||||
config = AutoRegressiveLMConfig(**config_kwargs)
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
model = AutoRegressiveLM(config).to(device=device)
|
||||
model.eval()
|
||||
|
||||
@@ -97,9 +86,10 @@ def test_model_forward(config_kwargs):
|
||||
|
||||
|
||||
@pytest.mark.parametrize("config_kwargs", CONFIGS)
|
||||
def test_model_forward_with_padding(config_kwargs):
|
||||
def test_model_forward_with_padding(config_kwargs, device):
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
|
||||
config = AutoRegressiveLMConfig(**config_kwargs)
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
model = AutoRegressiveLM(config).to(device=device)
|
||||
model.eval()
|
||||
|
||||
|
||||
+28
-29
@@ -249,17 +249,17 @@ def test_save_load_roundtrip():
|
||||
with torch.no_grad():
|
||||
out_src = model(x)["logits"].clone()
|
||||
|
||||
tmpdir = tempfile.mkdtemp()
|
||||
save_lora(model, tmpdir, cfg)
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
save_lora(model, tmpdir, cfg)
|
||||
|
||||
model2 = _make_model()
|
||||
model2.load_state_dict(model.state_dict(), strict=False)
|
||||
load_lora(model2, tmpdir)
|
||||
model2 = _make_model()
|
||||
model2.load_state_dict(model.state_dict(), strict=False)
|
||||
load_lora(model2, tmpdir)
|
||||
|
||||
with torch.no_grad():
|
||||
out_dst = model2(x)["logits"]
|
||||
with torch.no_grad():
|
||||
out_dst = model2(x)["logits"]
|
||||
|
||||
torch.testing.assert_close(out_src, out_dst)
|
||||
torch.testing.assert_close(out_src, out_dst)
|
||||
|
||||
|
||||
def test_save_after_merge_raises():
|
||||
@@ -271,13 +271,13 @@ def test_save_after_merge_raises():
|
||||
if isinstance(m, LoRALinear):
|
||||
m.lora_B.fill_(0.5)
|
||||
|
||||
tmpdir = tempfile.mkdtemp()
|
||||
save_lora(model, tmpdir, cfg)
|
||||
merge_lora(model)
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
save_lora(model, tmpdir, cfg)
|
||||
merge_lora(model)
|
||||
|
||||
tmpdir2 = tempfile.mkdtemp()
|
||||
with pytest.raises(RuntimeError, match="No LoRA parameters"):
|
||||
save_lora(model, tmpdir2, cfg)
|
||||
with tempfile.TemporaryDirectory() as tmpdir2:
|
||||
with pytest.raises(RuntimeError, match="No LoRA parameters"):
|
||||
save_lora(model, tmpdir2, cfg)
|
||||
|
||||
|
||||
def test_load_lora_on_already_injected():
|
||||
@@ -289,16 +289,15 @@ def test_load_lora_on_already_injected():
|
||||
if isinstance(m, LoRALinear):
|
||||
m.lora_B.fill_(0.5)
|
||||
|
||||
tmpdir = tempfile.mkdtemp()
|
||||
save_lora(model, tmpdir, LoRAConfig(r=4, alpha=8, target_modules=("q_proj",)))
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
save_lora(model, tmpdir, LoRAConfig(r=4, alpha=8, target_modules=("q_proj",)))
|
||||
|
||||
model2 = _make_model()
|
||||
model2.load_state_dict(model.state_dict(), strict=False)
|
||||
inject_lora(model2, r=4, alpha=8, target_modules={"q_proj"})
|
||||
model2 = _make_model()
|
||||
model2.load_state_dict(model.state_dict(), strict=False)
|
||||
inject_lora(model2, r=4, alpha=8, target_modules={"q_proj"})
|
||||
|
||||
# load onto already-injected model
|
||||
load_lora(model2, tmpdir)
|
||||
assert _get_lora_count(model2) > 0
|
||||
load_lora(model2, tmpdir)
|
||||
assert _get_lora_count(model2) > 0
|
||||
|
||||
|
||||
def test_load_lora_mismatched_r_raises():
|
||||
@@ -310,15 +309,15 @@ def test_load_lora_mismatched_r_raises():
|
||||
if isinstance(m, LoRALinear):
|
||||
m.lora_B.fill_(0.5)
|
||||
|
||||
tmpdir = tempfile.mkdtemp()
|
||||
save_lora(model, tmpdir, cfg)
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
save_lora(model, tmpdir, cfg)
|
||||
|
||||
model2 = _make_model()
|
||||
model2.load_state_dict(model.state_dict(), strict=False)
|
||||
inject_lora(model2, r=4, alpha=8, target_modules={"q_proj"})
|
||||
model2 = _make_model()
|
||||
model2.load_state_dict(model.state_dict(), strict=False)
|
||||
inject_lora(model2, r=4, alpha=8, target_modules={"q_proj"})
|
||||
|
||||
with pytest.raises(RuntimeError, match="size mismatch"):
|
||||
load_lora(model2, tmpdir) # strict=False, only lora keys
|
||||
with pytest.raises(RuntimeError, match="size mismatch"):
|
||||
load_lora(model2, tmpdir)
|
||||
|
||||
|
||||
def test_merge_preserves_output():
|
||||
|
||||
@@ -1,50 +1,18 @@
|
||||
import json
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
import pytest
|
||||
import safetensors.torch as st
|
||||
import torch
|
||||
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
from tests.helpers import TINY_CONFIG
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def transformer_test_env():
|
||||
test_dir = tempfile.mkdtemp(prefix="transformer_test_")
|
||||
config_path = os.path.join(test_dir, "config.json")
|
||||
def test_tie_weight_init(base_test_env):
|
||||
config_path = base_test_env["config_path"]
|
||||
|
||||
config = {
|
||||
"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,
|
||||
}
|
||||
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config, f)
|
||||
|
||||
yield {"test_dir": test_dir, "config_path": config_path, "config": config}
|
||||
|
||||
if os.path.exists(test_dir):
|
||||
try:
|
||||
for file in os.listdir(test_dir):
|
||||
os.remove(os.path.join(test_dir, file))
|
||||
os.rmdir(test_dir)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def test_tie_weight_init(transformer_test_env):
|
||||
config_path = transformer_test_env["config_path"]
|
||||
config_data = transformer_test_env["config"].copy()
|
||||
|
||||
# case 1: tie weight
|
||||
config_data = TINY_CONFIG.copy()
|
||||
config_data["tie_word_embeddings"] = True
|
||||
|
||||
with open(config_path, "w") as f:
|
||||
@@ -62,7 +30,6 @@ def test_tie_weight_init(transformer_test_env):
|
||||
assert torch.equal(model.lm_head.weight, model.embed_tokens.weight)
|
||||
assert not torch.equal(model.lm_head.weight, original_weight)
|
||||
|
||||
# case 2: not tie weight
|
||||
config_data["tie_word_embeddings"] = False
|
||||
|
||||
with open(config_path, "w") as f:
|
||||
@@ -81,13 +48,11 @@ def test_tie_weight_init(transformer_test_env):
|
||||
assert not torch.equal(model.lm_head.weight, original_weight)
|
||||
|
||||
|
||||
def test_model_save_load_with_tie_weight(transformer_test_env):
|
||||
test_dir = transformer_test_env["test_dir"]
|
||||
def test_model_save_load_with_tie_weight(base_test_env):
|
||||
test_dir = base_test_env["test_dir"]
|
||||
model_path = os.path.join(test_dir, "model.safetensors")
|
||||
|
||||
config_data = transformer_test_env["config"].copy()
|
||||
|
||||
# case 1: tie weight
|
||||
config_data = TINY_CONFIG.copy()
|
||||
config_data["tie_word_embeddings"] = True
|
||||
config_path = os.path.join(test_dir, "config.json")
|
||||
|
||||
@@ -107,7 +72,6 @@ def test_model_save_load_with_tie_weight(transformer_test_env):
|
||||
assert model.lm_head.weight.data_ptr() == model.embed_tokens.weight.data_ptr()
|
||||
assert "lm_head.weight" not in model.state_dict()
|
||||
|
||||
# case 2: not tie weight (form tie-weight state dict load)
|
||||
config_data["tie_word_embeddings"] = False
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_data, f)
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user