Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b99485f462 | ||
|
|
20041d7aa9 | ||
|
|
59248032dc | ||
|
|
ceadc34ea9 | ||
|
|
8ab5631446 | ||
|
|
99b5d2b2da | ||
|
|
021e6f3788 | ||
|
|
4e38183e86 | ||
|
|
4eeb23e2b3 | ||
|
|
ef8783b7e3 | ||
|
|
60d7ee614a | ||
|
|
f7a16efc9d | ||
|
|
a01e8bbe98 | ||
|
|
ccf728a1b7 | ||
|
|
f1b4b05d08 | ||
|
|
0c86c89af4 | ||
|
|
d7ac66fb73 | ||
|
|
a6e920fdb0 | ||
|
|
958df58f9d | ||
|
|
e0f102c4d9 | ||
|
|
5a942527b2 | ||
|
|
37a3036934 | ||
|
|
121a7bf8b4 | ||
|
|
a5678c9185 | ||
|
|
2c50b3cf37 | ||
|
|
eee7f54789 | ||
|
|
06eeeead79 | ||
|
|
e8ff7f5321 | ||
|
|
a6e1f26cd4 | ||
|
|
95c43368ae | ||
|
|
754624acf0 | ||
|
|
0b6a17330f | ||
|
|
74b9308883 | ||
|
|
e5f9b1a3a9 |
@@ -23,6 +23,7 @@ jobs:
|
||||
with:
|
||||
name: pure-wheel
|
||||
path: dist/*.whl
|
||||
if-no-files-found: error
|
||||
|
||||
build-cuda-linux:
|
||||
name: Build CUDA wheel (Linux)
|
||||
@@ -50,6 +51,7 @@ jobs:
|
||||
with:
|
||||
name: cuda-wheel-linux
|
||||
path: dist/*.whl
|
||||
if-no-files-found: error
|
||||
|
||||
release:
|
||||
name: Attach wheels to release
|
||||
@@ -58,14 +60,33 @@ jobs:
|
||||
permissions:
|
||||
contents: write
|
||||
steps:
|
||||
- uses: actions/download-artifact@v4
|
||||
- name: Download pure-Python wheel
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
pattern: "*-wheel"
|
||||
merge-multiple: true
|
||||
name: pure-wheel
|
||||
path: release-assets/pure
|
||||
|
||||
- name: Download CUDA wheel
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: cuda-wheel-linux
|
||||
path: release-assets/cuda
|
||||
|
||||
- name: Verify release assets
|
||||
shell: bash
|
||||
run: |
|
||||
set -euo pipefail
|
||||
pure_wheels=(release-assets/pure/*.whl)
|
||||
cuda_wheels=(release-assets/cuda/*.whl)
|
||||
test "${#pure_wheels[@]}" -eq 1
|
||||
test "${#cuda_wheels[@]}" -eq 1
|
||||
test "$(basename "${pure_wheels[0]}")" != "$(basename "${cuda_wheels[0]}")"
|
||||
|
||||
- name: Create release & upload assets
|
||||
uses: softprops/action-gh-release@v2
|
||||
with:
|
||||
files: ./*.whl
|
||||
files: |
|
||||
release-assets/pure/*.whl
|
||||
release-assets/cuda/*.whl
|
||||
tag_name: ${{ github.ref_name }}
|
||||
generate_release_notes: true
|
||||
|
||||
+239
-54
@@ -28,17 +28,17 @@ classDiagram
|
||||
|
||||
class AutoRegressiveLMConfig {
|
||||
+Optional[int] vocab_size
|
||||
+Optional[int] dim
|
||||
+Optional[int] n_layers
|
||||
+Optional[float] norm_eps
|
||||
+Optional[int] dim_ffn
|
||||
+Optional[bool] tie_weight
|
||||
+Optional[int] hidden_size
|
||||
+Optional[int] num_hidden_layers
|
||||
+Optional[float] rms_norm_eps
|
||||
+Optional[int] intermediate_size
|
||||
+Optional[bool] tie_word_embeddings
|
||||
+Optional[dict] rope_scaling
|
||||
+Optional[int] max_len
|
||||
+Optional[int] max_position_embeddings
|
||||
+Optional[float] rope_theta
|
||||
+str attn_type
|
||||
+Optional[int] n_heads
|
||||
+Optional[int] n_kv_heads
|
||||
+Optional[int] num_attention_heads
|
||||
+Optional[int] num_key_value_heads
|
||||
+Optional[bool] use_qk_norm
|
||||
+Optional[bool] use_gated_attention
|
||||
+Optional[int] kv_lora_rank
|
||||
@@ -53,15 +53,15 @@ classDiagram
|
||||
|
||||
class EncoderConfig {
|
||||
+Optional[int] vocab_size
|
||||
+Optional[int] dim
|
||||
+Optional[int] n_layers
|
||||
+Optional[float] norm_eps
|
||||
+Optional[int] dim_ffn
|
||||
+Optional[int] max_len
|
||||
+Optional[int] hidden_size
|
||||
+Optional[int] num_hidden_layers
|
||||
+Optional[float] rms_norm_eps
|
||||
+Optional[int] intermediate_size
|
||||
+Optional[int] max_position_embeddings
|
||||
+Optional[float] rope_theta
|
||||
+str attn_type
|
||||
+Optional[int] n_heads
|
||||
+Optional[int] n_kv_heads
|
||||
+Optional[int] num_attention_heads
|
||||
+Optional[int] num_key_value_heads
|
||||
+Optional[bool] use_qk_norm
|
||||
+str ffn_type
|
||||
+Optional[dict] rope_scaling
|
||||
@@ -141,6 +141,12 @@ classDiagram
|
||||
+int val_step
|
||||
+float neftune_alpha
|
||||
+str parallel_mode
|
||||
+int rollout_interval
|
||||
+float rollout_temperature
|
||||
+int rollout_top_k
|
||||
+float rollout_top_p
|
||||
+int rollout_max_tokens
|
||||
+Optional[Callable] reward_model_fn
|
||||
+dict executor_kwargs
|
||||
+dict extra_kwargs
|
||||
+validate()
|
||||
@@ -166,13 +172,6 @@ classDiagram
|
||||
+__getitem__(index) Dict
|
||||
}
|
||||
|
||||
class RecordDataset {
|
||||
+Optional[Callable] processor
|
||||
+load(load_path, storage_type)
|
||||
+__getitem__(index)
|
||||
+__len__()
|
||||
}
|
||||
|
||||
class DPODataset {
|
||||
+__getitem__(index) Dict
|
||||
}
|
||||
@@ -222,7 +221,12 @@ classDiagram
|
||||
+fetch_record(index, keys)
|
||||
}
|
||||
|
||||
class ResumableDistributedSampler {
|
||||
class JsonlSource {
|
||||
+Path path
|
||||
+load() List[dict]
|
||||
}
|
||||
|
||||
class RDSampler {
|
||||
+int epoch
|
||||
+int iter
|
||||
}
|
||||
@@ -385,19 +389,103 @@ classDiagram
|
||||
+forward(x) Tensor
|
||||
+set_neftune_alpha(alpha)
|
||||
}
|
||||
|
||||
class LoRAConfig {
|
||||
+int r
|
||||
+int alpha
|
||||
+tuple target_modules
|
||||
}
|
||||
|
||||
class LoRALinear {
|
||||
+Linear weight
|
||||
+Parameter lora_A, lora_B
|
||||
+forward(x) Tensor
|
||||
+merge()
|
||||
}
|
||||
}
|
||||
|
||||
namespace preprocessing {
|
||||
class SectionRenderer {
|
||||
+process_sections(item, sections, config, tokenizer) Tuple
|
||||
+process_list_field(item, sections, config, tokenizer) Tuple
|
||||
}
|
||||
|
||||
class BaseMaskBuilder {
|
||||
<<abstract>>
|
||||
+build(item, config, tokenizer) Optional[dict]
|
||||
}
|
||||
|
||||
class SectionedMaskBuilder {
|
||||
class SingleOutputMaskBuilder {
|
||||
+SectionRenderer renderer
|
||||
+build(item, config, tokenizer) Optional[dict]
|
||||
+_build_single(item, config, tokenizer) Optional[dict]
|
||||
+_build_multi(item, sources_spec, config, tokenizer) Optional[dict]
|
||||
}
|
||||
|
||||
class MultiOutputMaskBuilder {
|
||||
+SectionRenderer renderer
|
||||
+build(item, config, tokenizer) Optional[dict]
|
||||
}
|
||||
|
||||
class SectionedMaskBuilder {
|
||||
+build(item, config, tokenizer) Optional[dict]
|
||||
}
|
||||
|
||||
class PackingStrategy {
|
||||
<<abstract>>
|
||||
+apply(keys, max_packed_len, truncation_mode) Dict
|
||||
}
|
||||
|
||||
class PackingStrategyFactory {
|
||||
+create(name, *args, **kwargs) PackingStrategy
|
||||
}
|
||||
|
||||
class SimplePacking {
|
||||
+apply(keys, max_packed_len, truncation_mode) Dict
|
||||
}
|
||||
|
||||
class BFDPacking {
|
||||
+apply(keys, max_packed_len, truncation_mode) Dict
|
||||
}
|
||||
|
||||
class BFDSplitPacking {
|
||||
+apply(keys, max_packed_len, truncation_mode) Dict
|
||||
}
|
||||
|
||||
class PositionIdStrategy {
|
||||
<<abstract>>
|
||||
+generate(sequences) List[int]
|
||||
}
|
||||
|
||||
class PositionIdStrategyFactory {
|
||||
+create(name, *args, **kwargs) PositionIdStrategy
|
||||
}
|
||||
|
||||
class NoPositionId {
|
||||
+generate(sequences) List[int]
|
||||
}
|
||||
|
||||
class DocResetPositionId {
|
||||
+generate(sequences) List[int]
|
||||
}
|
||||
|
||||
class ContinuousPositionId {
|
||||
+generate(sequences) List[int]
|
||||
}
|
||||
|
||||
class StoreWriter {
|
||||
<<abstract>>
|
||||
+save(output_dir, domain, shard_idx, tensors)
|
||||
}
|
||||
|
||||
class StoreWriterFactory {
|
||||
+create(name, *args, **kwargs) StoreWriter
|
||||
}
|
||||
|
||||
class BinWriter {
|
||||
+save(output_dir, domain, shard_idx, tensors)
|
||||
}
|
||||
|
||||
class H5Writer {
|
||||
+save(output_dir, domain, shard_idx, tensors)
|
||||
}
|
||||
|
||||
class Pipeline {
|
||||
@@ -497,7 +585,7 @@ classDiagram
|
||||
|
||||
class TrainContextBuilder {
|
||||
+TrainConfig config
|
||||
+with_resume_dir(resume_dir) TrainContextBuilder
|
||||
+with_param_path(param_path, resume) TrainContextBuilder
|
||||
+build() TrainContext
|
||||
}
|
||||
|
||||
@@ -544,6 +632,32 @@ classDiagram
|
||||
+sync_old_model()
|
||||
}
|
||||
|
||||
class RawRollout {
|
||||
+Tensor prompts
|
||||
+Tensor responses
|
||||
+Tensor response_mask
|
||||
+Tensor logprobs_old
|
||||
}
|
||||
|
||||
class RolloutResult {
|
||||
+Tensor rewards
|
||||
}
|
||||
|
||||
class BaseRewardModel {
|
||||
<<abstract>>
|
||||
+score(prompts, responses) Tensor
|
||||
}
|
||||
|
||||
class RolloutGenerator {
|
||||
+generate(batch) RawRollout
|
||||
}
|
||||
|
||||
class RolloutRunner {
|
||||
+step()
|
||||
+clear_cache()
|
||||
+__call__(batch) Tuple[RolloutResult, bool]
|
||||
}
|
||||
|
||||
class BaseScheduler {
|
||||
+get_lr() List[float]
|
||||
+step()
|
||||
@@ -857,12 +971,21 @@ classDiagram
|
||||
+apply(logits, filter_value) Tensor
|
||||
}
|
||||
|
||||
class FrequencyPenaltyStrategy {
|
||||
+float penalty
|
||||
+apply(logits, filter_value, input_ids, input_mask) Tensor
|
||||
}
|
||||
|
||||
class SamplingPipeline {
|
||||
+List[BaseSamplingStrategy] strategies
|
||||
+apply(logits, filter_value) Tensor
|
||||
+sample(logits, filter_value) Tensor
|
||||
}
|
||||
|
||||
class StreamDecoder {
|
||||
+push(token_id) str
|
||||
}
|
||||
|
||||
class GenerateResult {
|
||||
+List[Tuple[int, str]] tokens
|
||||
+List[str] results
|
||||
@@ -881,6 +1004,17 @@ classDiagram
|
||||
+Optional[str] tool_call_id
|
||||
}
|
||||
|
||||
class FunctionDef {
|
||||
+str name
|
||||
+Optional[str] description
|
||||
+Optional[Dict] parameters
|
||||
}
|
||||
|
||||
class ToolDef {
|
||||
+str type
|
||||
+FunctionDef function
|
||||
}
|
||||
|
||||
class ChatCompletionRequest {
|
||||
+str model
|
||||
+List[ChatMessage] messages
|
||||
@@ -969,9 +1103,20 @@ classDiagram
|
||||
+str yielded
|
||||
}
|
||||
|
||||
class get_app {
|
||||
<<module>>
|
||||
+get_app() FastAPI
|
||||
class BaseToolParser {
|
||||
<<abstract>>
|
||||
+feed(body, current_token_ids, delta_token_ids) List[Dict]
|
||||
+parse_complete(body) Optional[Dict]
|
||||
+has_tool_calls (property) bool
|
||||
}
|
||||
|
||||
class ToolParserFactory {
|
||||
+create(name, *args, **kwargs) BaseToolParser
|
||||
}
|
||||
|
||||
class SimpleJsonToolParser {
|
||||
+feed(body, current_token_ids, delta_token_ids) List[Dict]
|
||||
+parse_complete(body) Optional[Dict]
|
||||
}
|
||||
}
|
||||
|
||||
@@ -994,14 +1139,17 @@ classDiagram
|
||||
}
|
||||
|
||||
namespace parallel {
|
||||
class setup {
|
||||
<<module>>
|
||||
+spawn_parallel_fn(func, world_size, backend, master_addr, master_port, device_type, start_method, **kwargs)
|
||||
+setup_parallel(rank, world_size, backend, master_addr, master_port, device_type) contextmanager
|
||||
+get_current_device() str
|
||||
+get_world_size() int
|
||||
+get_rank() int
|
||||
+only_on_rank(rank, sync=False) decorator
|
||||
class LaunchStrategy {
|
||||
<<abstract>>
|
||||
+launch(func, **kwargs)
|
||||
}
|
||||
|
||||
class TorchrunStrategy {
|
||||
+launch(func, **kwargs)
|
||||
}
|
||||
|
||||
class LocalStrategy {
|
||||
+launch(func, **kwargs)
|
||||
}
|
||||
|
||||
class GradientState {
|
||||
@@ -1030,7 +1178,7 @@ classDiagram
|
||||
|
||||
class BaseExecutor {
|
||||
+GradientState gradient_state
|
||||
+prepare(model, optimizer, dataloader, scheduler) tuple
|
||||
+prepare(model_fn, optimizer_fn, scheduler_fn, before_wrap) tuple
|
||||
+accumulate(model) context manager
|
||||
+backward(loss)
|
||||
+unwrap_model(model) dict
|
||||
@@ -1052,6 +1200,12 @@ classDiagram
|
||||
+unwrap_model(model) dict
|
||||
}
|
||||
|
||||
class FSDP2Executor {
|
||||
-_prepare_model(model) nn.Module
|
||||
-_no_sync(model) context manager
|
||||
+unwrap_model(model) dict
|
||||
}
|
||||
|
||||
class ExecutorFactory {
|
||||
+Dict _entries
|
||||
+register(name) decorator
|
||||
@@ -1104,9 +1258,8 @@ classDiagram
|
||||
TrainCallback <|-- MetricCallback
|
||||
BaseDataset <|-- SEQDataset
|
||||
BaseDataset <|-- SFTDataset
|
||||
BaseDataset <|-- RecordDataset
|
||||
RecordDataset <|-- DPODataset
|
||||
RecordDataset <|-- GRPODataset
|
||||
BaseDataset <|-- DPODataset
|
||||
BaseDataset <|-- GRPODataset
|
||||
Store <|-- H5Store
|
||||
Store <|-- MmapStore
|
||||
Store <|-- JsonlStore
|
||||
@@ -1119,6 +1272,7 @@ classDiagram
|
||||
BaseSamplingStrategy <|-- TemperatureStrategy
|
||||
BaseSamplingStrategy <|-- TopKStrategy
|
||||
BaseSamplingStrategy <|-- TopPStrategy
|
||||
BaseSamplingStrategy <|-- FrequencyPenaltyStrategy
|
||||
ParallelModel <|-- RowParallelLinear
|
||||
ParallelModel <|-- ColumnParallelLinear
|
||||
AutoModel <|-- AutoRegressiveLM
|
||||
@@ -1142,12 +1296,31 @@ classDiagram
|
||||
BaseFactory <|-- ExecutorFactory
|
||||
BaseFactory <|-- ConfigFactory
|
||||
BaseFactory <|-- MaskBuilderFactory
|
||||
BaseFactory <|-- PackingStrategyFactory
|
||||
BaseFactory <|-- PositionIdStrategyFactory
|
||||
BaseFactory <|-- StoreWriterFactory
|
||||
BaseFactory <|-- ToolParserFactory
|
||||
BaseExecutor <|-- NoneExecutor
|
||||
BaseExecutor <|-- DDPExecutor
|
||||
BaseExecutor <|-- FSDPExecutor
|
||||
BaseExecutor <|-- FSDP2Executor
|
||||
ResponseBuilder <|-- OpenAIResponseBuilder
|
||||
ResponseBuilder <|-- AnthropicResponseBuilder
|
||||
BaseToolParser <|-- SimpleJsonToolParser
|
||||
BaseMaskBuilder <|-- SectionedMaskBuilder
|
||||
BaseMaskBuilder <|-- SingleOutputMaskBuilder
|
||||
BaseMaskBuilder <|-- MultiOutputMaskBuilder
|
||||
PackingStrategy <|-- SimplePacking
|
||||
PackingStrategy <|-- BFDPacking
|
||||
BFDPacking <|-- BFDSplitPacking
|
||||
PositionIdStrategy <|-- NoPositionId
|
||||
PositionIdStrategy <|-- DocResetPositionId
|
||||
PositionIdStrategy <|-- ContinuousPositionId
|
||||
StoreWriter <|-- BinWriter
|
||||
StoreWriter <|-- H5Writer
|
||||
RawRollout <|-- RolloutResult
|
||||
LaunchStrategy <|-- TorchrunStrategy
|
||||
LaunchStrategy <|-- LocalStrategy
|
||||
KVCache <|-- PageCache
|
||||
KVCache <|-- ContiguousCache
|
||||
CacheView <|-- PageCacheView
|
||||
@@ -1169,6 +1342,8 @@ classDiagram
|
||||
EmbeddingEncoder *-- Embedding
|
||||
DecoderBlock *-- RMSNorm
|
||||
ChatCompletionRequest *-- ChatMessage
|
||||
ChatCompletionRequest *-- ToolDef
|
||||
ToolDef *-- FunctionDef
|
||||
MessagesRequest *-- AnthropicMessage
|
||||
BaseExecutor *-- GradientState
|
||||
AccumOptimizer o-- GradientState
|
||||
@@ -1191,6 +1366,9 @@ classDiagram
|
||||
Pipeline o-- PipelineConfig
|
||||
Pipeline o-- BaseMaskBuilder
|
||||
Pipeline o-- AutoTokenizer
|
||||
Pipeline o-- PackingStrategy
|
||||
Pipeline o-- PositionIdStrategy
|
||||
Pipeline o-- StoreWriter
|
||||
TokenizeTransform o-- AutoTokenizer
|
||||
TokenizeTransform o-- BaseMaskBuilder
|
||||
|
||||
@@ -1198,6 +1376,9 @@ classDiagram
|
||||
TrainConfig ..> BaseStrategy : selects
|
||||
PipelineConfig ..> MaskBuilderFactory : selects
|
||||
MaskBuilderFactory ..> BaseMaskBuilder : creates
|
||||
PackingStrategyFactory ..> PackingStrategy : creates
|
||||
PositionIdStrategyFactory ..> PositionIdStrategy : creates
|
||||
StoreWriterFactory ..> StoreWriter : creates
|
||||
StrategyFactory ..> BaseStrategy : creates
|
||||
SchedulerFactory ..> BaseScheduler : creates
|
||||
DatasetFactory ..> BaseDataset : creates
|
||||
@@ -1216,12 +1397,13 @@ classDiagram
|
||||
ExecutorFactory ..> NoneExecutor : creates
|
||||
ExecutorFactory ..> DDPExecutor : creates
|
||||
ExecutorFactory ..> FSDPExecutor : creates
|
||||
ExecutorFactory ..> FSDP2Executor : creates
|
||||
ToolParserFactory ..> BaseToolParser : creates
|
||||
TrainContextBuilder ..> ExecutorFactory : creates
|
||||
Trainer ..> TrainContextBuilder : uses
|
||||
TrainContextBuilder ..> TrainContext : creates
|
||||
Trainer ..> Functions : spawns
|
||||
TrainContextBuilder ..> StrategyFactory : uses
|
||||
TrainContextBuilder ..> ResumableDistributedSampler : creates
|
||||
TrainContextBuilder ..> RDSampler : creates
|
||||
Checkpoint ..> Checkpoint : serializes
|
||||
CheckpointCallback ..> Checkpoint : creates
|
||||
PageCache ..> PageCacheView : binds
|
||||
@@ -1232,6 +1414,9 @@ classDiagram
|
||||
AnthropicResponseBuilder ..> MessagesRequest : receives
|
||||
ProtocolHandler ..> StopChecker : creates
|
||||
ProtocolHandler ..> GenContext : creates
|
||||
RolloutGenerator ..> InferenceScheduler : uses
|
||||
RolloutRunner ..> RolloutGenerator : uses
|
||||
RolloutRunner ..> BaseRewardModel : uses
|
||||
|
||||
%% --- Association (general usage) ---
|
||||
Trainer --> TrainConfig
|
||||
@@ -1253,14 +1438,14 @@ classDiagram
|
||||
| Module | Components | Description |
|
||||
|--------|------------|-------------|
|
||||
| **astrai.config** | BaseConfig, BaseModelConfig, AutoRegressiveLMConfig, EncoderConfig, ConfigFactory, TrainConfig, PipelineConfig, InputConfig, ProcessingConfig, OutputConfig | Configuration management (to_dict/from_dict, to_file/from_file) |
|
||||
| **astrai.preprocessing** | BaseMaskBuilder, MaskBuilderFactory, SectionedMaskBuilder, SingleOutputMaskBuilder, MultiOutputMaskBuilder, Pipeline, TokenizeTransform, filter_by_length, PackingStrategy, PackingStrategyFactory, plan_bfd, PositionIdStrategy, PositionIdStrategyFactory, StoreWriter, StoreWriterFactory, core (shared helpers) | Declarative JSON-driven data preprocessing |
|
||||
| **astrai.dataset** | BaseDataset–RecordDataset–DPO/GRPODataset, SEQDataset, SFTDataset, Store, Streamable, Recordable, H5Store, MmapStore, JsonlStore, StoreFactory, ResumableDistributedSampler, DatasetFactory | Dataset loading and management |
|
||||
| **astrai.preprocessing** | SectionRenderer, BaseMaskBuilder, MaskBuilderFactory, SectionedMaskBuilder, SingleOutputMaskBuilder, MultiOutputMaskBuilder, Pipeline, TokenizeTransform, PackingStrategy, PackingStrategyFactory, SimplePacking, BFDPacking, BFDSplitPacking, PositionIdStrategy, PositionIdStrategyFactory, NoPositionId, DocResetPositionId, ContinuousPositionId, StoreWriter, StoreWriterFactory, BinWriter, H5Writer | Declarative JSON-driven data preprocessing |
|
||||
| **astrai.dataset** | BaseDataset, SEQDataset, SFTDataset, DPODataset, GRPODataset, Store, Streamable, Recordable, H5Store, MmapStore, JsonlSource, JsonlStore, StoreFactory, RDSampler, DatasetFactory | Dataset loading and management |
|
||||
| **astrai.serialization** | Checkpoint | Model serialization |
|
||||
| **astrai.model** | AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, RotaryEmbedding, Embedding | Neural network model |
|
||||
| **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 | Training workflow |
|
||||
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, KVCache–ContiguousCache/PageCache, CacheView–ContiguousCacheView/PageCacheView, Allocator–Storage, Task, TaskManager, TaskStatus, GenerationRequest, GenerateResult, BaseSamplingStrategy–SamplingPipeline, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, ChatMessage–MessagesRequest, app | Inference service |
|
||||
| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, FSDPExecutor, GradientState, AccumOptimizer, AccumScheduler, ParallelModel, RowParallelLinear, ColumnParallelLinear | Distributed parallel & gradient accumulation |
|
||||
| **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.factory** | BaseFactory | Component registration |
|
||||
| **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers |
|
||||
|
||||
@@ -1268,7 +1453,7 @@ classDiagram
|
||||
|
||||
| Pattern | Classes | Purpose |
|
||||
|---------|---------|---------|
|
||||
| **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory`, `MaskBuilderFactory`, `StoreWriterFactory`, `PackingStrategyFactory`, `PositionIdStrategyFactory` | Decorator-based component creation |
|
||||
| **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory`, `MaskBuilderFactory`, `StoreWriterFactory`, `PackingStrategyFactory`, `PositionIdStrategyFactory`, `ToolParserFactory` | Decorator-based component creation |
|
||||
| **Registry** | `BaseFactory` | Component registration |
|
||||
| **Strategy** | `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy` | Training strategy switching |
|
||||
| **Strategy (Sampling)** | `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations |
|
||||
@@ -1277,7 +1462,7 @@ 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` | Gradient accumulation & model distribution |
|
||||
| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor`, `FSDP2Executor` | 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 |
|
||||
@@ -1287,7 +1472,7 @@ 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`
|
||||
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`
|
||||
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`
|
||||
@@ -1296,4 +1481,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-19
|
||||
> Document Update Time: 2026-07-20
|
||||
|
||||
@@ -85,7 +85,7 @@ All backends normalise tensors into `Store._data[Dict[str, List[Tensor]]]` + `St
|
||||
```
|
||||
DatasetFactory.load(train_type, load_path, window_size, stride=None,
|
||||
storage_type=None, tokenizer_path=None,
|
||||
max_len=2048, store=None)
|
||||
max_position_embeddings=2048, store=None)
|
||||
→ BaseDataset.load(load_path, storage_type=None)
|
||||
→ detect_format(load_path)
|
||||
→ StoreFactory.create(storage_type)
|
||||
|
||||
@@ -32,7 +32,7 @@ ContiguousCache (simple contiguous per-slot cache)
|
||||
├── ContiguousCacheView bundles k/v tensors + slot indices for attention layers
|
||||
```
|
||||
|
||||
Created by default when no cache is passed to `InferenceScheduler`. Each task occupies a fixed slot of `[max_seq_len, n_kv_heads, head_dim]`. Simple and efficient for small-to-medium batch sizes.
|
||||
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.
|
||||
|
||||
### PageCache (paged with prefix sharing)
|
||||
|
||||
@@ -42,7 +42,7 @@ PageCache (paged KV cache with prefix sharing, alternative)
|
||||
│ ├── 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 (n_layers × n_pages × page_size × n_kv_heads × head_dim)
|
||||
├── 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
|
||||
```
|
||||
|
||||
|
||||
+21
-8
@@ -13,7 +13,7 @@
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--train_type` | Training type (`seq`, `sft`, `dpo`, `grpo`) | required |
|
||||
| `--train_type` | Training type (`seq`, `sft`, `dpo`, `grpo`, `online_grpo`, `online_dpo`) | required |
|
||||
| `--data_root_path` | Dataset root directory | required |
|
||||
| `--param_path` | Model parameters or checkpoint path | required |
|
||||
| `--n_epoch` | Total training epochs | 1 |
|
||||
@@ -26,7 +26,7 @@
|
||||
|-----------|-------------|---------|
|
||||
| `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 |
|
||||
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
|
||||
| `--max_grad_norm` | Maximum gradient norm for clipping (None disables) | None |
|
||||
| `--max_grad_norm` | Maximum gradient norm for clipping (None disables) | 1.0 |
|
||||
|
||||
### Optimizer (MuonMix)
|
||||
|
||||
@@ -44,7 +44,7 @@ Combined optimizer: matrix parameters via **Muon**, non-matrix via **AdamW** (`f
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--window_size` | Max input sequence length | model config `max_len` |
|
||||
| `--window_size` | Max input sequence length | model config `max_position_embeddings` |
|
||||
| `--stride` | Stride for sliding window over sequences | None |
|
||||
| `--random_seed` | Random seed for reproducibility | 3407 |
|
||||
| `--num_workers` | DataLoader worker processes | 4 |
|
||||
@@ -100,18 +100,31 @@ Combined optimizer: matrix parameters via **Muon**, non-matrix via **AdamW** (`f
|
||||
| `--group_size` | GRPO group size | 4 | `grpo` |
|
||||
| `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo` |
|
||||
| `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 | `grpo` |
|
||||
| `--grpo_sync_interval` | GRPO ref_model sync interval (steps) | 200 | `grpo` |
|
||||
| `--neftune_alpha` | NEFTune noise alpha (0=disabled, typical: 5.0) | 0.0 | `sft` |
|
||||
|
||||
### Online Rollout
|
||||
|
||||
These options apply to `online_grpo` and `online_dpo`. Online strategies require
|
||||
a `BaseRewardModel` factory in `TrainConfig`; `train.py` does not currently
|
||||
provide a command-line option for configuring one.
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--rollout_interval` | Optimizer steps between rollout refreshes | 512 |
|
||||
| `--rollout_temperature` | Rollout sampling temperature | 0.7 |
|
||||
| `--rollout_top_k` | Rollout top-k filtering (`0` disables) | 0 |
|
||||
| `--rollout_top_p` | Rollout nucleus sampling threshold | 0.9 |
|
||||
| `--rollout_max_tokens` | Maximum generated tokens per response | 1024 |
|
||||
|
||||
### Scheduler
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--schedule_type` | LR scheduler type (`cosine`, `sgdr`, `wsd`) | cosine |
|
||||
| `--min_rate` | Minimum LR as fraction of base LR | None (scheduler default: 0.01) |
|
||||
| `--min_rate` | Minimum LR as fraction of base LR | None (scheduler default: 0.05 for cosine/SGDR, 0.0 for WSD) |
|
||||
| `--cycle_length` | SGDR first cycle length in steps | None (total_steps - warmup_steps) |
|
||||
| `--t_mult` | SGDR cycle length multiplier per restart | 2 |
|
||||
| `--stable_steps` | WSD stable plateau steps | None (required for wsd) |
|
||||
| `--stable_steps` | WSD stable plateau steps | None (80% of post-warmup steps) |
|
||||
| `--decay_steps` | WSD decay steps | None (total_steps - warmup_steps - stable_steps) |
|
||||
|
||||
### Usage Example
|
||||
@@ -173,7 +186,7 @@ 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 |
|
||||
| `--max_tokens` | int | model config `max_len` | Maximum tokens to generate |
|
||||
| `--max_tokens` | int | model config `max_position_embeddings` | Maximum tokens to generate |
|
||||
|
||||
Usage:
|
||||
```bash
|
||||
@@ -201,4 +214,4 @@ See [Preprocessing Guide](preprocessing.md) for config file format and examples.
|
||||
|
||||
---
|
||||
|
||||
> Document Update Time: 2026-07-19
|
||||
> Document Update Time: 2026-07-20
|
||||
|
||||
+20
-10
@@ -6,7 +6,7 @@
|
||||
- [Causal Mask](#causal-mask)
|
||||
- [Rotary Position Embedding (RoPE)](#rotary-position-embedding-rope)
|
||||
- [Training Loop](#training-loop)
|
||||
- [Strategies](#strategies) — SEQ, SFT, DPO, GRPO
|
||||
- [Strategies](#strategies) — SEQ, SFT, DPO, GRPO, online rollout
|
||||
- [LR Schedulers](#lr-schedulers)
|
||||
- [Gradient Checkpointing](#gradient-checkpointing)
|
||||
- [Checkpoint](#checkpoint)
|
||||
@@ -146,6 +146,19 @@ Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`. External sync of `ol
|
||||
|
||||
Keys: `prompts`, `responses`, `masks`, `rewards`.
|
||||
|
||||
### Online Rollout
|
||||
|
||||
`online_grpo` and `online_dpo` use the respective GRPO and DPO strategies with
|
||||
a `RolloutRunner`. The runner renders prompts through the tokenizer chat
|
||||
template, generates grouped responses through `InferenceScheduler`, then scores
|
||||
them with a `BaseRewardModel`. It refreshes cached rollouts every
|
||||
`rollout_interval` optimizer steps. `online_grpo` synchronizes `old_model` when
|
||||
a fresh rollout is produced.
|
||||
|
||||
Online strategies require `TrainConfig.reward_model_fn`. `train.py` exposes the
|
||||
rollout sampling parameters but does not yet offer a CLI argument for the reward
|
||||
model factory.
|
||||
|
||||
## LR Schedulers
|
||||
|
||||
| Type | Class | Description |
|
||||
@@ -162,6 +175,7 @@ Trades compute for memory by recomputing activations during backward pass. Speci
|
||||
|
||||
```python
|
||||
from astrai.model.components.decoder_block import DecoderBlock
|
||||
|
||||
config = TrainConfig(..., gradient_checkpointing_modules=[DecoderBlock])
|
||||
```
|
||||
|
||||
@@ -181,18 +195,14 @@ Model config (`context.model_config`) saved into `config.json` during training v
|
||||
## TrainContextBuilder (Builder Pattern)
|
||||
|
||||
```python
|
||||
context = (
|
||||
TrainContextBuilder(config)
|
||||
.with_resume_dir(resume_dir)
|
||||
.build()
|
||||
)
|
||||
context = TrainContextBuilder(config).with_param_path(param_path, resume=True).build()
|
||||
# Returns TrainContext with model, strategy, optimizer, scheduler, dataloader, checkpoint
|
||||
```
|
||||
|
||||
- Loads checkpoint weights if provided
|
||||
- Loads checkpoint weights before the model is wrapped
|
||||
- Creates executor via `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)`
|
||||
- Calls `executor.prepare(model, optimizer, dataloader, scheduler)` for model distribution (e.g. DDP) + gradient accumulation wrappers
|
||||
- Creates `ResumableDistributedSampler` for shuffle+resume
|
||||
- Calls `executor.prepare(model_fn, optimizer_fn, scheduler_fn, before_wrap=...)`; the executor creates, wraps, then builds the optimizer and scheduler for the wrapped model
|
||||
- Creates `RDSampler` for shuffle+resume
|
||||
- Builds strategy via `StrategyFactory.create(train_type, model, device, **kwargs)`
|
||||
|
||||
## Training CLI
|
||||
@@ -222,4 +232,4 @@ nohup python scripts/tools/train.py \
|
||||
|
||||
Full parameter reference at [params.md](params.md).
|
||||
|
||||
> Document Update Time: 2026-07-19
|
||||
> Document Update Time: 2026-07-20
|
||||
|
||||
+1
-1
@@ -1,4 +1,4 @@
|
||||
__version__ = "1.3.10"
|
||||
__version__ = "1.3.11"
|
||||
__author__ = "ViperEkura"
|
||||
|
||||
from astrai.config import (
|
||||
|
||||
@@ -29,19 +29,19 @@ class AutoRegressiveLMConfig(BaseModelConfig):
|
||||
"""Configuration for autoregressive language model."""
|
||||
|
||||
vocab_size: Optional[int] = None
|
||||
dim: Optional[int] = None
|
||||
n_layers: Optional[int] = None
|
||||
norm_eps: Optional[float] = None
|
||||
dim_ffn: Optional[int] = None
|
||||
tie_weight: Optional[bool] = None
|
||||
hidden_size: Optional[int] = None
|
||||
num_hidden_layers: Optional[int] = None
|
||||
rms_norm_eps: Optional[float] = None
|
||||
intermediate_size: Optional[int] = None
|
||||
tie_word_embeddings: Optional[bool] = None
|
||||
|
||||
max_len: Optional[int] = None
|
||||
max_position_embeddings: Optional[int] = None
|
||||
rope_theta: Optional[float] = None
|
||||
rope_scaling: Optional[dict] = None
|
||||
|
||||
attn_type: str = "gqa"
|
||||
n_heads: Optional[int] = None
|
||||
n_kv_heads: Optional[int] = None
|
||||
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
|
||||
|
||||
@@ -62,18 +62,18 @@ class EncoderConfig(BaseModelConfig):
|
||||
"""Configuration for embedding encoder model."""
|
||||
|
||||
vocab_size: Optional[int] = None
|
||||
dim: Optional[int] = None
|
||||
n_layers: Optional[int] = None
|
||||
norm_eps: Optional[float] = None
|
||||
dim_ffn: 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_len: Optional[int] = None
|
||||
max_position_embeddings: Optional[int] = None
|
||||
rope_theta: Optional[float] = None
|
||||
rope_scaling: Optional[dict] = None
|
||||
|
||||
attn_type: str = "gqa"
|
||||
n_heads: Optional[int] = None
|
||||
n_kv_heads: Optional[int] = None
|
||||
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
|
||||
|
||||
|
||||
@@ -45,6 +45,8 @@ class ProcessingConfig(BaseConfig):
|
||||
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.
|
||||
|
||||
@@ -65,6 +67,7 @@ class ProcessingConfig(BaseConfig):
|
||||
min_chars: int = 50
|
||||
max_chars: int = 2_000_000
|
||||
max_items: Optional[int] = None
|
||||
batch_size: int = 256
|
||||
packing_strategy: str = "simple"
|
||||
max_packed_len: int = 8192
|
||||
truncation_mode: str = "keep_start"
|
||||
|
||||
@@ -38,7 +38,7 @@ class TrainConfig(BaseConfig):
|
||||
default=1, metadata={"help": "Number of iterations between steps."}
|
||||
)
|
||||
max_grad_norm: Optional[float] = field(
|
||||
default=None,
|
||||
default=1.0,
|
||||
metadata={"help": "Maximum gradient norm. None disables clipping."},
|
||||
)
|
||||
gradient_checkpointing_modules: List[str] = field(
|
||||
@@ -138,6 +138,32 @@ class TrainConfig(BaseConfig):
|
||||
metadata={"help": "NEFTune noise alpha (0=disabled, typical: 5.0)."},
|
||||
)
|
||||
|
||||
# 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)."
|
||||
},
|
||||
)
|
||||
|
||||
executor_kwargs: Dict[str, Any] = field(
|
||||
default_factory=dict,
|
||||
metadata={"help": "Extra kwargs passed to ExecutorFactory.create()."},
|
||||
|
||||
@@ -190,7 +190,8 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
|
||||
- rewards: [G]
|
||||
|
||||
Output:
|
||||
- prompts: [B, P_max]
|
||||
- prompts: [B, P_max], left-padded
|
||||
- prompt_mask: [B, P_max]
|
||||
- responses: [B, G, R_max]
|
||||
- masks: [B, G, R_max]
|
||||
- rewards: [B, G]
|
||||
@@ -201,13 +202,15 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
|
||||
R_max = max(r.size(0) for b in batch for r in b["responses"])
|
||||
|
||||
prompts = torch.zeros(B, P_max, dtype=torch.long)
|
||||
prompt_mask = torch.zeros(B, P_max, dtype=torch.bool)
|
||||
responses = torch.zeros(B, G, R_max, dtype=torch.long)
|
||||
masks = torch.zeros(B, G, R_max, dtype=torch.bool)
|
||||
rewards = torch.zeros(B, G, dtype=torch.float32)
|
||||
|
||||
for i, b in enumerate(batch):
|
||||
p_len = b["prompts"].size(0)
|
||||
prompts[i, :p_len] = b["prompts"]
|
||||
prompts[i, -p_len:] = b["prompts"]
|
||||
prompt_mask[i, -p_len:] = True
|
||||
rewards[i, : b["rewards"].size(0)] = b["rewards"]
|
||||
for g in range(min(G, len(b["responses"]))):
|
||||
r_len = b["responses"][g].size(0)
|
||||
@@ -217,6 +220,7 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
|
||||
|
||||
return {
|
||||
"prompts": prompts,
|
||||
"prompt_mask": prompt_mask,
|
||||
"responses": responses,
|
||||
"masks": masks,
|
||||
"rewards": rewards,
|
||||
@@ -346,7 +350,15 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
|
||||
if processor is not None:
|
||||
store.load(load_path, processor=processor, **kwargs)
|
||||
else:
|
||||
store.load(load_path, **kwargs)
|
||||
load_kwargs = dict(kwargs)
|
||||
if (
|
||||
tokenizer_path is not None
|
||||
and storage_type == "jsonl"
|
||||
and train_type in ("seq", "sft")
|
||||
and "tokenizer_path" not in load_kwargs
|
||||
):
|
||||
load_kwargs["tokenizer_path"] = tokenizer_path
|
||||
store.load(load_path, **load_kwargs)
|
||||
|
||||
return cls.create(train_type, store=store)
|
||||
|
||||
|
||||
+42
-14
@@ -56,6 +56,7 @@ from typing import Callable, Dict, List, Optional, Tuple, Union
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.config.preprocess_config import PipelineConfig
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.preprocessing.transform import TokenizeTransform
|
||||
from astrai.serialization import (
|
||||
@@ -545,18 +546,29 @@ class JsonlSource:
|
||||
|
||||
@StoreFactory.register("jsonl")
|
||||
class JsonlStore(Store, Streamable, Recordable):
|
||||
"""JSONL reader with two tokenisation modes.
|
||||
"""JSONL reader with eager/lazy tokenisation modes.
|
||||
|
||||
A JSONL dataset is a ``.jsonl`` file or a directory of ``*.jsonl``
|
||||
files plus (optionally) a ``dataset_config.json`` describing the
|
||||
tokenization pipeline.
|
||||
|
||||
Two modes, selected at :meth:`load` time:
|
||||
Three ways to supply an eager transform (first match wins):
|
||||
|
||||
- **Eager** (default): applies a :class:`TokenizeTransform` to every
|
||||
record at load time and registers per-key tensors via
|
||||
``_normalize``. Both ``fetch`` (stream) and ``fetch_record``
|
||||
(record) work.
|
||||
- **Explicit** (``transform=``): caller-built
|
||||
:class:`TokenizeTransform` applied eagerly.
|
||||
- **Config file**: ``dataset_config.json`` alongside the ``*.jsonl``
|
||||
files — loaded via :meth:`TokenizeTransform.from_config_file`.
|
||||
- **Default messages** (``tokenizer_path=`` given, no config file):
|
||||
a built-in chatml config that tokenises the ``messages`` field,
|
||||
masking every role except ``assistant`` (loss on assistant only).
|
||||
Lets SFT/SEQ train straight from a chat-style JSONL directory
|
||||
without a hand-written config.
|
||||
|
||||
Two tokenisation modes, selected at :meth:`load` time:
|
||||
|
||||
- **Eager** (default): applies the transform to every record at load
|
||||
time and registers per-key tensors via ``_normalize``. Both
|
||||
``fetch`` (stream) and ``fetch_record`` (record) work.
|
||||
- **Lazy** (``processor=fn`` passed): keeps raw records and defers
|
||||
tokenisation to ``fetch_record``. Only record access works —
|
||||
``len(store)`` returns ``num_records``; stream primitives raise.
|
||||
@@ -565,6 +577,16 @@ class JsonlStore(Store, Streamable, Recordable):
|
||||
CONFIG_NAME = "dataset_config.json"
|
||||
segments_are_records = True
|
||||
|
||||
_DEFAULT_MESSAGES_CONFIG = {
|
||||
"version": 1,
|
||||
"input": {
|
||||
"sections": [{"field": "messages", "action": "$role", "template": True}]
|
||||
},
|
||||
"mask": {"system": "mask", "user": "mask", "assistant": "train"},
|
||||
"mask_default": "mask",
|
||||
"output": {"position_ids_mode": "doc_reset"},
|
||||
}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
window_size: int = 0,
|
||||
@@ -587,14 +609,20 @@ class JsonlStore(Store, Streamable, Recordable):
|
||||
if transform is None:
|
||||
root = Path(path)
|
||||
config_path = root / self.CONFIG_NAME if root.is_dir() else None
|
||||
if config_path is None or not config_path.exists():
|
||||
raise FileNotFoundError(
|
||||
f"JSONL dataset config not found. Expected "
|
||||
f"{self.CONFIG_NAME} alongside *.jsonl files, pass an "
|
||||
f"explicit transform, or pass processor= for lazy "
|
||||
f"on-the-fly tokenisation."
|
||||
)
|
||||
transform = TokenizeTransform.from_config_file(str(config_path))
|
||||
if config_path is not None and config_path.exists():
|
||||
transform = TokenizeTransform.from_config_file(str(config_path))
|
||||
else:
|
||||
tokenizer_path = kwargs.get("tokenizer_path")
|
||||
if not tokenizer_path:
|
||||
raise FileNotFoundError(
|
||||
f"JSONL dataset config not found. Expected "
|
||||
f"{self.CONFIG_NAME} alongside *.jsonl files, pass an "
|
||||
f"explicit transform, pass processor= for lazy "
|
||||
f"on-the-fly tokenisation, or pass tokenizer_path= to "
|
||||
f"use the built-in messages config."
|
||||
)
|
||||
config = PipelineConfig.from_dict(self._DEFAULT_MESSAGES_CONFIG)
|
||||
transform = TokenizeTransform(config, tokenizer_path)
|
||||
|
||||
transformed = transform.apply(records)
|
||||
self._normalize(transformed)
|
||||
|
||||
@@ -18,12 +18,13 @@ when available, otherwise falls back to ``torch.nn.functional.scaled_dot_product
|
||||
"""
|
||||
|
||||
from astrai.extension.loader import KERNEL_NAMES, is_available
|
||||
from astrai.extension.ops import attn_decode, attn_paged_decode, attn_prefill
|
||||
from astrai.extension.ops import attention, attn_decode, attn_paged_decode, attn_prefill
|
||||
|
||||
__all__ = [
|
||||
"attn_decode",
|
||||
"attn_paged_decode",
|
||||
"attn_prefill",
|
||||
"attention",
|
||||
"is_available",
|
||||
"KERNEL_NAMES",
|
||||
]
|
||||
|
||||
@@ -244,3 +244,55 @@ def attn_paged_decode(
|
||||
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
|
||||
|
||||
+4
-1
@@ -67,7 +67,10 @@ class BaseFactory(ABC, Generic[T]):
|
||||
if _get_origin(orig_base) is BaseFactory:
|
||||
(arg,) = _get_args(orig_base)
|
||||
cls._entries = {}
|
||||
cls._component_base = _resolve_type(arg, cls)
|
||||
try:
|
||||
cls._component_base = _resolve_type(arg, cls)
|
||||
except Exception:
|
||||
cls._component_base = None
|
||||
return
|
||||
|
||||
@classmethod
|
||||
|
||||
@@ -435,24 +435,13 @@ class ContiguousCacheView(CacheView):
|
||||
pos = self._write_positions
|
||||
self._cache.k[layer_id, indices, pos] = k.squeeze(1)
|
||||
self._cache.v[layer_id, indices, pos] = v.squeeze(1)
|
||||
for s, p in zip(indices.tolist(), pos.tolist()):
|
||||
cur = self._cache._slot_len.get(s, 0)
|
||||
if p + 1 > cur:
|
||||
self._cache._slot_len[s] = p + 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
|
||||
new_len = start_pos + seq_len
|
||||
for s in indices.tolist():
|
||||
cur = self._cache._slot_len.get(s, 0)
|
||||
if new_len > cur:
|
||||
self._cache._slot_len[s] = new_len
|
||||
|
||||
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
|
||||
max_len = max(
|
||||
self._cache._slot_len.get(int(s), 0) for s in self._batch_indices.tolist()
|
||||
)
|
||||
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]
|
||||
@@ -528,6 +517,9 @@ class ContiguousCache(KVCache):
|
||||
) -> 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
|
||||
)
|
||||
|
||||
@@ -43,19 +43,40 @@ class Executor:
|
||||
)
|
||||
|
||||
task_ids = [t.task_id for t in tasks]
|
||||
position_ids = (
|
||||
torch.arange(start_pos, prompt_len, dtype=torch.long, device=self.device)
|
||||
.unsqueeze(0)
|
||||
.expand(batch_sz, -1)
|
||||
)
|
||||
input_mask = position_ids.unsqueeze(-1) >= torch.arange(
|
||||
prompt_len, device=self.device
|
||||
)
|
||||
|
||||
with torch.inference_mode():
|
||||
self.model(
|
||||
input_ids,
|
||||
position_ids=torch.arange(
|
||||
start_pos, prompt_len, dtype=torch.long, device=self.device
|
||||
)
|
||||
.unsqueeze(0)
|
||||
.expand(batch_sz, -1),
|
||||
input_mask=input_mask,
|
||||
position_ids=position_ids,
|
||||
paged_cache=self.kv_cache.bind_tasks(task_ids, prompt_len, self.device),
|
||||
)
|
||||
|
||||
def execute_decode(self, tasks: List[Task]) -> List[int]:
|
||||
def execute_decode(
|
||||
self, tasks: List[Task], return_logprobs: bool = False
|
||||
) -> List[int]:
|
||||
"""Decode next token for each task.
|
||||
|
||||
Args:
|
||||
return_logprobs: When ``True``, also record (and return)
|
||||
the log-probability of each sampled token under the
|
||||
post-strategy sampling distribution. The logprob is
|
||||
appended to ``task.output_logprobs`` and the return
|
||||
list becomes ``List[Tuple[int, float]]``.
|
||||
|
||||
Returns:
|
||||
``List[int]`` of sampled token IDs, or
|
||||
``List[Tuple[int, float]]`` of ``(token_id, logprob)`` when
|
||||
``return_logprobs`` is ``True``.
|
||||
"""
|
||||
if not tasks:
|
||||
return []
|
||||
|
||||
@@ -68,7 +89,10 @@ class Executor:
|
||||
position_ids = torch.tensor(
|
||||
[t.next_pos for t in tasks], dtype=torch.long, device=self.device
|
||||
)
|
||||
total_len = position_ids.max().item() + 1
|
||||
total_len = max(t.next_pos for t in tasks) + 1
|
||||
input_mask = position_ids[:, None, None] >= torch.arange(
|
||||
total_len, device=self.device
|
||||
)
|
||||
|
||||
task_ids = [t.task_id for t in tasks]
|
||||
|
||||
@@ -80,32 +104,30 @@ class Executor:
|
||||
)
|
||||
|
||||
history_lists = []
|
||||
mask_lists = []
|
||||
history_lens = []
|
||||
for t in tasks:
|
||||
window = t.rep_window
|
||||
prompt_part = t.prompt_ids[-window:]
|
||||
ids = prompt_part + t.output_ids
|
||||
history_lists.append(ids)
|
||||
mask_lists.append([True] * len(ids))
|
||||
history_lens.append(len(ids))
|
||||
|
||||
max_len = max(len(h) for h in history_lists)
|
||||
max_len = max(history_lens) if history_lens else 0
|
||||
padded_ids = torch.zeros(
|
||||
len(tasks), max_len, dtype=torch.long, device=self.device
|
||||
)
|
||||
padded_mask = torch.zeros(
|
||||
len(tasks), max_len, dtype=torch.bool, device=self.device
|
||||
)
|
||||
for i, (h, m) in enumerate(zip(history_lists, mask_lists)):
|
||||
padded_ids[i, : len(h)] = torch.tensor(
|
||||
h, dtype=torch.long, device=self.device
|
||||
)
|
||||
padded_mask[i, : len(m)] = torch.tensor(
|
||||
m, dtype=torch.bool, device=self.device
|
||||
)
|
||||
for i, h in enumerate(history_lists):
|
||||
L = history_lens[i]
|
||||
padded_ids[i, :L] = torch.as_tensor(h, dtype=torch.long, device=self.device)
|
||||
padded_mask[i, :L] = True
|
||||
|
||||
with torch.inference_mode():
|
||||
outputs = self.model(
|
||||
input_ids.unsqueeze(1),
|
||||
input_mask=input_mask,
|
||||
paged_cache=self.kv_cache.bind_tasks(
|
||||
task_ids,
|
||||
total_len,
|
||||
@@ -116,6 +138,23 @@ class Executor:
|
||||
)
|
||||
logits = outputs["logits"][:, -1, :]
|
||||
|
||||
if return_logprobs:
|
||||
tokens, logprobs = sample(
|
||||
logits,
|
||||
temperature=temperatures,
|
||||
top_k=top_ks,
|
||||
top_p=top_ps,
|
||||
frequency_penalty=freq_penalties,
|
||||
input_ids=padded_ids,
|
||||
input_mask=padded_mask,
|
||||
return_logprobs=True,
|
||||
)
|
||||
tokens_list = tokens.tolist()
|
||||
logprobs_list = logprobs.tolist()
|
||||
for t, lp in zip(tasks, logprobs_list):
|
||||
t.output_logprobs.append(float(lp))
|
||||
return list(zip(tokens_list, logprobs_list))
|
||||
|
||||
return sample(
|
||||
logits,
|
||||
temperature=temperatures,
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import logging
|
||||
import threading
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
@@ -31,26 +32,26 @@ class InferenceScheduler:
|
||||
|
||||
if max_seq_len is not None:
|
||||
self.max_seq_len = max_seq_len
|
||||
elif config.max_len is not None:
|
||||
self.max_seq_len = config.max_len
|
||||
elif config.max_position_embeddings is not None:
|
||||
self.max_seq_len = config.max_position_embeddings
|
||||
else:
|
||||
raise ValueError(
|
||||
"max_seq_len must be provided either as argument "
|
||||
"or in model config (config.max_len)"
|
||||
"or in model config (config.max_position_embeddings)"
|
||||
)
|
||||
self.device = device or next(model.parameters()).device
|
||||
self.dtype = dtype or next(model.parameters()).dtype
|
||||
|
||||
head_dim = config.dim // config.n_heads
|
||||
head_dim = config.hidden_size // config.num_attention_heads
|
||||
|
||||
if cache is not None:
|
||||
self._cache = cache
|
||||
else:
|
||||
self._cache = ContiguousCache(
|
||||
config.n_layers,
|
||||
config.num_hidden_layers,
|
||||
max_batch_size,
|
||||
self.max_seq_len,
|
||||
config.n_kv_heads,
|
||||
config.num_key_value_heads,
|
||||
head_dim,
|
||||
self.device,
|
||||
self.dtype,
|
||||
@@ -194,6 +195,117 @@ class InferenceScheduler:
|
||||
self._cache.task_free(task.task_id)
|
||||
for task in self._task_mgr.get_waiting_tasks():
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
self._cache.task_free(task.task_id)
|
||||
self._task_mgr.clear_queues()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def run_batch(
|
||||
self,
|
||||
prompt_ids_list: List[List[int]],
|
||||
*,
|
||||
max_tokens: Optional[int] = None,
|
||||
temperature: float = 1.0,
|
||||
top_p: float = 1.0,
|
||||
top_k: int = 50,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
return_logprobs: bool = False,
|
||||
) -> List[List[int]]:
|
||||
"""Synchronous batch generation without the scheduler thread.
|
||||
|
||||
Accepts already-tokenized prompts (no string round-trip) and runs
|
||||
prefill + decode to completion on the calling thread. Designed for
|
||||
RL rollout, where logprobs of the behaviour policy must be collected
|
||||
alongside generated tokens.
|
||||
|
||||
Args:
|
||||
prompt_ids_list: ``B`` prompts, each a list of token IDs.
|
||||
max_tokens: Maximum tokens to generate per prompt. ``None``
|
||||
uses ``self.max_seq_len - len(prompt_ids)``.
|
||||
temperature/top_p/top_k/frequency_penalty/rep_window: Sampling
|
||||
parameters (uniform across the batch).
|
||||
return_logprobs: If ``True``, return ``(token_ids, logprobs)``
|
||||
tuples per prompt (logprobs aligned 1-to-1 with token_ids).
|
||||
|
||||
Returns:
|
||||
``List[List[int]]`` of generated token IDs per prompt, or —
|
||||
when ``return_logprobs`` is ``True`` —
|
||||
``List[Tuple[List[int], List[float]]]``.
|
||||
"""
|
||||
stop_ids = self._task_mgr.tokenizer.stop_ids
|
||||
cache = self._cache
|
||||
seq_cap = self.max_seq_len
|
||||
|
||||
tasks: List[Task] = []
|
||||
for ids in prompt_ids_list:
|
||||
if len(ids) >= seq_cap:
|
||||
tasks.append(None)
|
||||
continue
|
||||
t_max = max_tokens
|
||||
if t_max is None:
|
||||
t_max = seq_cap - len(ids)
|
||||
else:
|
||||
t_max = min(t_max, seq_cap - len(ids))
|
||||
task = Task(
|
||||
task_id=f"batch_{uuid.uuid4().hex[:8]}",
|
||||
prompt_ids=list(ids),
|
||||
max_tokens=t_max,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
frequency_penalty=frequency_penalty,
|
||||
rep_window=rep_window,
|
||||
)
|
||||
if not cache.task_alloc(task.task_id, task.prompt_ids):
|
||||
tasks.append(None)
|
||||
continue
|
||||
task.input_tokens = len(task.prompt_ids)
|
||||
tasks.append(task)
|
||||
|
||||
try:
|
||||
live = [t for t in tasks if t is not None]
|
||||
prefill_groups: Dict[Tuple[int, int], List[Task]] = {}
|
||||
for t in live:
|
||||
key = (len(t.prompt_ids), cache.task_cached(t.task_id))
|
||||
prefill_groups.setdefault(key, []).append(t)
|
||||
for (prompt_len, start_pos), group in prefill_groups.items():
|
||||
self._executor.execute_prefill(group, prompt_len, start_pos)
|
||||
|
||||
while live:
|
||||
valid: List[Task] = []
|
||||
for t in sorted(live, key=lambda x: x.task_id):
|
||||
if cache.task_extend(t.task_id, t.next_pos):
|
||||
valid.append(t)
|
||||
else:
|
||||
t.status = TaskStatus.ABORTED
|
||||
if not valid:
|
||||
break
|
||||
|
||||
step_out = self._executor.execute_decode(
|
||||
valid, return_logprobs=return_logprobs
|
||||
)
|
||||
if return_logprobs:
|
||||
for t, (ntok, _lp) in zip(valid, step_out):
|
||||
t.output_ids.append(ntok)
|
||||
t.output_tokens += 1
|
||||
else:
|
||||
for t, ntok in zip(valid, step_out):
|
||||
t.output_ids.append(ntok)
|
||||
t.output_tokens += 1
|
||||
|
||||
live = [t for t in valid if not t.is_finished(stop_ids)]
|
||||
finally:
|
||||
for t in tasks:
|
||||
if t is not None:
|
||||
cache.task_free(t.task_id)
|
||||
|
||||
results: List[Any] = []
|
||||
for t in tasks:
|
||||
if t is None:
|
||||
results.append(([], []) if return_logprobs else [])
|
||||
elif return_logprobs:
|
||||
results.append((list(t.output_ids), list(t.output_logprobs)))
|
||||
else:
|
||||
results.append(list(t.output_ids))
|
||||
return results
|
||||
|
||||
@@ -81,6 +81,7 @@ class Task:
|
||||
|
||||
self.status = TaskStatus.PENDING
|
||||
self.output_ids: List[int] = []
|
||||
self.output_logprobs: List[float] = []
|
||||
self.input_tokens: int = 0
|
||||
self.output_tokens: int = 0
|
||||
self.arrival_time = time.time()
|
||||
|
||||
+49
-17
@@ -276,7 +276,8 @@ class SamplingPipeline(BaseSamplingStrategy):
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
return_logprobs: bool = False,
|
||||
):
|
||||
"""Apply strategies then sample (softmax + multinomial).
|
||||
|
||||
Short-circuits to ``argmax`` when temperature is exactly 0
|
||||
@@ -286,21 +287,41 @@ class SamplingPipeline(BaseSamplingStrategy):
|
||||
logits: Raw logits ``[batch, vocab_size]``.
|
||||
input_ids: Previously generated token IDs ``[batch, seq_len]``.
|
||||
input_mask: Boolean mask for ``input_ids`` padding.
|
||||
return_logprobs: If ``True``, return ``(tokens, logprobs)``
|
||||
where ``logprobs[i]`` is the log-probability of
|
||||
``tokens[i]`` under the (post-strategy) sampling
|
||||
distribution.
|
||||
|
||||
Returns:
|
||||
Sampled token IDs ``[batch]``.
|
||||
Sampled token IDs ``[batch]``, or — when ``return_logprobs``
|
||||
is ``True`` — a ``(token_ids, chosen_logprobs)`` tuple.
|
||||
"""
|
||||
for s in self.strategies:
|
||||
if isinstance(s, TemperatureStrategy) and self._is_greedy(s.temperature):
|
||||
return logits.argmax(dim=-1)
|
||||
break
|
||||
if self._is_greedy_pipeline():
|
||||
tokens = logits.argmax(dim=-1)
|
||||
if not return_logprobs:
|
||||
return tokens
|
||||
log_probs = torch.log_softmax(logits.float(), dim=-1)
|
||||
chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1)
|
||||
return tokens, chosen
|
||||
|
||||
return torch.multinomial(
|
||||
torch.softmax(
|
||||
self.apply(logits, filter_value, input_ids, input_mask), dim=-1
|
||||
),
|
||||
num_samples=1,
|
||||
transformed = self.apply(logits, filter_value, input_ids, input_mask)
|
||||
log_probs = torch.log_softmax(transformed.float(), dim=-1)
|
||||
tokens = torch.multinomial(
|
||||
torch.softmax(transformed, dim=-1), num_samples=1
|
||||
).squeeze(-1)
|
||||
if not return_logprobs:
|
||||
return tokens
|
||||
chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1)
|
||||
return tokens, chosen
|
||||
|
||||
def _is_greedy_pipeline(self) -> bool:
|
||||
"""True if the first strategy is greedy temperature (temp=0)."""
|
||||
if not self.strategies:
|
||||
return False
|
||||
first = self.strategies[0]
|
||||
return isinstance(first, TemperatureStrategy) and self._is_greedy(
|
||||
first.temperature
|
||||
)
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
@@ -313,10 +334,11 @@ def sample(
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
filter_value: float = -float("inf"),
|
||||
) -> Tensor:
|
||||
return_logprobs: bool = False,
|
||||
):
|
||||
"""Apply sampling strategies then sample (softmax + multinomial).
|
||||
|
||||
Shortcut for ``SamplingPipeline(...).sample(logits)``.
|
||||
Shortcut for ``SamplingPipeline(...).sample(logits, return_logprobs=)``.
|
||||
|
||||
When **temperature** is exactly 0 (scalar or single-element tensor)
|
||||
the function short-circuits to ``argmax`` for deterministic decode.
|
||||
@@ -327,12 +349,16 @@ def sample(
|
||||
(0.0 disables, range -2.0~2.0).
|
||||
input_ids: Previously generated token IDs ``[batch, seq_len]``.
|
||||
input_mask: Boolean mask for ``input_ids`` padding.
|
||||
return_logprobs: If ``True``, also return the log-probability
|
||||
of each sampled token under the (post-strategy) sampling
|
||||
distribution — useful for RL rollout (PPO/GRPO importance
|
||||
ratios).
|
||||
|
||||
Returns:
|
||||
Sampled token IDs ``[batch]``.
|
||||
Sampled token IDs ``[batch]``, or — when ``return_logprobs`` is
|
||||
``True`` — a ``(token_ids, chosen_logprobs)`` tuple where
|
||||
``chosen_logprobs`` has shape ``[batch]``.
|
||||
"""
|
||||
if SamplingPipeline._is_greedy(temperature):
|
||||
return logits.argmax(dim=-1)
|
||||
return SamplingPipeline(
|
||||
[
|
||||
TemperatureStrategy(temperature),
|
||||
@@ -340,4 +366,10 @@ def sample(
|
||||
TopPStrategy(top_p),
|
||||
FrequencyPenaltyStrategy(frequency_penalty),
|
||||
]
|
||||
).sample(logits, filter_value, input_ids, input_mask)
|
||||
).sample(
|
||||
logits,
|
||||
filter_value=filter_value,
|
||||
input_ids=input_ids,
|
||||
input_mask=input_mask,
|
||||
return_logprobs=return_logprobs,
|
||||
)
|
||||
|
||||
@@ -76,9 +76,8 @@ class GQA(nn.Module):
|
||||
rotary_emb: Tensor,
|
||||
attn_mask: Tensor = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
is_causal = attn_mask is None
|
||||
|
||||
q = self._split_heads(self.q_proj(x), self.n_heads)
|
||||
k = self._split_heads(self.k_proj(x), self.n_kv_heads)
|
||||
v = self._split_heads(self.v_proj(x), self.n_kv_heads)
|
||||
@@ -163,9 +162,9 @@ class MLA(nn.Module):
|
||||
rotary_emb: Tensor,
|
||||
attn_mask: Tensor = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
bsz, seq_len, _ = x.size()
|
||||
is_causal = attn_mask is None
|
||||
|
||||
q = self.q_proj(x)
|
||||
q = q.view(bsz, seq_len, self.n_heads, self.head_dim)
|
||||
|
||||
@@ -14,10 +14,18 @@ class DecoderBlock(nn.Module):
|
||||
def __init__(self, config, layer_id: int):
|
||||
super().__init__()
|
||||
cfg = asdict(config)
|
||||
cfg["down_init_std"] = 0.02 / (2 * config.n_layers) ** 0.5
|
||||
cfg.update(
|
||||
dim=config.hidden_size,
|
||||
dim_ffn=config.intermediate_size,
|
||||
n_layers=config.num_hidden_layers,
|
||||
n_heads=config.num_attention_heads,
|
||||
n_kv_heads=config.num_key_value_heads,
|
||||
norm_eps=config.rms_norm_eps,
|
||||
down_init_std=0.02 / (2 * config.num_hidden_layers) ** 0.5,
|
||||
)
|
||||
self.attention = AttnFactory.create(config.attn_type, **cfg, layer_id=layer_id)
|
||||
self.input_norm = RMSNorm(config.dim, config.norm_eps)
|
||||
self.post_attention_norm = RMSNorm(config.dim, config.norm_eps)
|
||||
self.input_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
||||
self.post_attention_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
||||
self.mlp = FFNFactory.create(config.ffn_type, **cfg)
|
||||
|
||||
def forward(
|
||||
@@ -26,12 +34,14 @@ class DecoderBlock(nn.Module):
|
||||
rotary_emb: Tensor,
|
||||
attention_mask: Optional[Tensor] = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
attn_output = self.attention(
|
||||
self.input_norm(x),
|
||||
rotary_emb,
|
||||
attention_mask,
|
||||
paged_cache,
|
||||
is_causal,
|
||||
)
|
||||
x = attn_output + x
|
||||
x = self.mlp(self.post_attention_norm(x)) + x
|
||||
|
||||
@@ -39,8 +39,12 @@ class LoRALinear(nn.Module):
|
||||
|
||||
self.r = r
|
||||
self.scaling = alpha / r
|
||||
self.lora_A = nn.Parameter(torch.randn(r, self.weight.shape[1]) / r)
|
||||
self.lora_B = nn.Parameter(torch.zeros(self.weight.shape[0], r))
|
||||
device = self.weight.device
|
||||
dtype = self.weight.dtype
|
||||
lora_a = torch.randn(r, self.weight.shape[1], device=device, dtype=dtype) / r
|
||||
lora_b = torch.zeros(self.weight.shape[0], r, device=device, dtype=dtype)
|
||||
self.lora_A = nn.Parameter(lora_a)
|
||||
self.lora_B = nn.Parameter(lora_b)
|
||||
self._merged = False
|
||||
|
||||
def forward(self, x):
|
||||
|
||||
+15
-7
@@ -18,20 +18,28 @@ class EmbeddingEncoder(AutoModel):
|
||||
def __init__(self, config: EncoderConfig):
|
||||
super().__init__(config)
|
||||
self.config = config
|
||||
rope_dim = config.dim // config.n_heads
|
||||
rope_dim = config.hidden_size // config.num_attention_heads
|
||||
rope_base = config.rope_theta if config.rope_theta is not None else 10000
|
||||
self.rotary_embedding = RotaryEmbedding(
|
||||
rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling
|
||||
rope_dim,
|
||||
config.max_position_embeddings,
|
||||
rope_base,
|
||||
rope_scaling=config.rope_scaling,
|
||||
)
|
||||
self.embed_tokens = Embedding(
|
||||
config.vocab_size, config.dim, neftune_alpha=config.neftune_alpha
|
||||
config.vocab_size,
|
||||
config.hidden_size,
|
||||
neftune_alpha=config.neftune_alpha,
|
||||
)
|
||||
|
||||
self.layers = nn.ModuleList(
|
||||
[DecoderBlock(config, layer_id) for layer_id in range(config.n_layers)]
|
||||
[
|
||||
DecoderBlock(config, layer_id)
|
||||
for layer_id in range(config.num_hidden_layers)
|
||||
]
|
||||
)
|
||||
|
||||
self.norm = RMSNorm(config.dim, config.norm_eps)
|
||||
self.norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
||||
|
||||
self.pooling_type = config.pooling_type or "mean"
|
||||
self.normalize_embeddings = config.normalize_embeddings or False
|
||||
@@ -59,10 +67,10 @@ class EmbeddingEncoder(AutoModel):
|
||||
x = self.embed_tokens(input_ids)
|
||||
|
||||
rotary_emb = self.rotary_embedding(x, position_ids)
|
||||
attn_mask = process_attention_mask(x, position_ids, input_mask, is_causal=False)
|
||||
attn_mask = process_attention_mask(input_mask)
|
||||
|
||||
for layer in self.layers:
|
||||
x = layer(x, rotary_emb, attn_mask, paged_cache=None)
|
||||
x = layer(x, rotary_emb, attn_mask)
|
||||
|
||||
hidden_states = self.norm(x)
|
||||
|
||||
|
||||
+27
-35
@@ -15,32 +15,15 @@ from astrai.model.components.rope import RotaryEmbedding
|
||||
|
||||
|
||||
def process_attention_mask(
|
||||
input_tensor: Tensor,
|
||||
position_ids: Optional[Tensor],
|
||||
input_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
input_mask: Optional[Tensor],
|
||||
) -> Optional[Tensor]:
|
||||
if position_ids is None:
|
||||
return None
|
||||
if input_mask is not None and input_mask.dim() > 2:
|
||||
return input_mask
|
||||
|
||||
device = input_tensor.device
|
||||
B = input_tensor.size(0)
|
||||
T = position_ids.max().item() + 1
|
||||
|
||||
if input_mask is None:
|
||||
if position_ids.min().item() == 0 and is_causal:
|
||||
return None
|
||||
attend = torch.ones(B, 1, T, dtype=torch.bool, device=device)
|
||||
else:
|
||||
attend = input_mask[:, :T].to(device=device, dtype=torch.bool).unsqueeze(1)
|
||||
|
||||
if is_causal:
|
||||
causal = position_ids.unsqueeze(-1) >= torch.arange(T, device=device)
|
||||
attend = attend & causal
|
||||
|
||||
return attend.unsqueeze(1)
|
||||
return None
|
||||
if input_mask.dim() == 2:
|
||||
return input_mask[:, None, None, :]
|
||||
if input_mask.dim() == 3:
|
||||
return input_mask[:, None, :, :]
|
||||
return input_mask
|
||||
|
||||
|
||||
@AutoModel.register("autoregressive_lm")
|
||||
@@ -53,24 +36,32 @@ class AutoRegressiveLM(AutoModel):
|
||||
rope_dim = (
|
||||
config.qk_rope_head_dim
|
||||
if config.attn_type == "mla"
|
||||
else config.dim // config.n_heads
|
||||
else config.hidden_size // config.num_attention_heads
|
||||
)
|
||||
rope_base = config.rope_theta if config.rope_theta is not None else 10000
|
||||
self.rotary_embedding = RotaryEmbedding(
|
||||
rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling
|
||||
rope_dim,
|
||||
config.max_position_embeddings,
|
||||
rope_base,
|
||||
rope_scaling=config.rope_scaling,
|
||||
)
|
||||
self.embed_tokens = Embedding(
|
||||
config.vocab_size, config.dim, neftune_alpha=config.neftune_alpha
|
||||
config.vocab_size,
|
||||
config.hidden_size,
|
||||
neftune_alpha=config.neftune_alpha,
|
||||
)
|
||||
|
||||
self.layers = nn.ModuleList(
|
||||
[DecoderBlock(config, layer_id) for layer_id in range(config.n_layers)]
|
||||
[
|
||||
DecoderBlock(config, layer_id)
|
||||
for layer_id in range(config.num_hidden_layers)
|
||||
]
|
||||
)
|
||||
|
||||
self.norm = RMSNorm(config.dim, config.norm_eps)
|
||||
self.lm_head = Linear(config.dim, config.vocab_size)
|
||||
self.norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
||||
self.lm_head = Linear(config.hidden_size, config.vocab_size)
|
||||
|
||||
if self.config.tie_weight is True:
|
||||
if self.config.tie_word_embeddings is True:
|
||||
self.lm_head.weight = self.embed_tokens.weight
|
||||
|
||||
self.apply(self._init_weights)
|
||||
@@ -85,7 +76,7 @@ class AutoRegressiveLM(AutoModel):
|
||||
|
||||
state_dict = dict(state_dict)
|
||||
|
||||
if self.config.tie_weight is True:
|
||||
if self.config.tie_word_embeddings is True:
|
||||
# same tensor for embed and lm_head
|
||||
if embed_key in state_dict:
|
||||
state_dict[lm_head_key] = state_dict[embed_key]
|
||||
@@ -101,7 +92,7 @@ class AutoRegressiveLM(AutoModel):
|
||||
destination=destination, prefix=prefix, keep_vars=keep_vars
|
||||
)
|
||||
|
||||
if self.config.tie_weight is True:
|
||||
if self.config.tie_word_embeddings is True:
|
||||
lm_head_key = prefix + "lm_head.weight"
|
||||
if lm_head_key in state_dict:
|
||||
del state_dict[lm_head_key]
|
||||
@@ -119,10 +110,11 @@ class AutoRegressiveLM(AutoModel):
|
||||
|
||||
x = self.embed_tokens(input_ids)
|
||||
rotary_emb = self.rotary_embedding(x, position_ids)
|
||||
attn_mask = process_attention_mask(x, position_ids, input_mask, is_causal=True)
|
||||
attn_mask = process_attention_mask(input_mask)
|
||||
use_sdpa_causal_mask = attn_mask is None
|
||||
|
||||
for layer in self.layers:
|
||||
x = layer(x, rotary_emb, attn_mask, paged_cache)
|
||||
x = layer(x, rotary_emb, attn_mask, paged_cache, use_sdpa_causal_mask)
|
||||
|
||||
hidden_states = self.norm(x)
|
||||
logits = self.lm_head(hidden_states)
|
||||
|
||||
@@ -4,6 +4,7 @@ from astrai.parallel.executor import (
|
||||
BaseExecutor,
|
||||
DDPExecutor,
|
||||
ExecutorFactory,
|
||||
FSDP2Executor,
|
||||
FSDPExecutor,
|
||||
GradientState,
|
||||
NoneExecutor,
|
||||
@@ -35,4 +36,5 @@ __all__ = [
|
||||
"NoneExecutor",
|
||||
"DDPExecutor",
|
||||
"FSDPExecutor",
|
||||
"FSDP2Executor",
|
||||
]
|
||||
|
||||
+121
-25
@@ -4,17 +4,22 @@ import contextlib
|
||||
import logging
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
from typing import Optional, Tuple
|
||||
from typing import Any, Callable, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
from torch.distributed.fsdp import FullStateDictConfig, StateDictType
|
||||
from torch.distributed.fsdp import (
|
||||
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
|
||||
from torch.optim.lr_scheduler import LRScheduler
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.parallel.setup import get_rank, get_world_size
|
||||
@@ -86,19 +91,25 @@ class BaseExecutor:
|
||||
|
||||
def prepare(
|
||||
self,
|
||||
model: nn.Module,
|
||||
optimizer: Optional[Optimizer] = None,
|
||||
dataloader: Optional[DataLoader] = None,
|
||||
scheduler: Optional[LRScheduler] = None,
|
||||
) -> Tuple[
|
||||
nn.Module, Optional[Optimizer], Optional[DataLoader], Optional[LRScheduler]
|
||||
]:
|
||||
model_fn: Callable[[], nn.Module],
|
||||
optimizer_fn: Optional[Callable[[nn.Module], Optimizer]] = None,
|
||||
scheduler_fn: Optional[Callable[[Optimizer], LRScheduler]] = None,
|
||||
before_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 optimizer is not None:
|
||||
optimizer = None
|
||||
scheduler = None
|
||||
if optimizer_fn is not None:
|
||||
optimizer = optimizer_fn(model)
|
||||
if scheduler_fn is not None:
|
||||
scheduler = scheduler_fn(optimizer)
|
||||
optimizer = AccumOptimizer(optimizer, self.gradient_state)
|
||||
if scheduler is not None:
|
||||
scheduler = AccumScheduler(scheduler, self.gradient_state)
|
||||
return model, optimizer, dataloader, scheduler
|
||||
if scheduler is not None:
|
||||
scheduler = AccumScheduler(scheduler, self.gradient_state)
|
||||
return model, optimizer, scheduler
|
||||
|
||||
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
||||
return model
|
||||
@@ -148,14 +159,7 @@ class BaseExecutor:
|
||||
def grad_accum_steps(self) -> int:
|
||||
return self.gradient_state.num_steps
|
||||
|
||||
def clip_grad_norm(self, model: nn.Module, max_norm: Optional[float]) -> float:
|
||||
if max_norm is None:
|
||||
total_norm = torch.norm(
|
||||
torch.stack(
|
||||
[p.grad.norm(2) for p in model.parameters() if p.grad is not None]
|
||||
)
|
||||
)
|
||||
return total_norm.item()
|
||||
def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
|
||||
total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)
|
||||
if isinstance(total_norm, torch.Tensor):
|
||||
return total_norm.item()
|
||||
@@ -289,9 +293,7 @@ class FSDPExecutor(BaseExecutor):
|
||||
return model.no_sync()
|
||||
return contextlib.nullcontext()
|
||||
|
||||
def clip_grad_norm(self, model: nn.Module, max_norm: Optional[float]) -> float:
|
||||
if max_norm is None:
|
||||
return super().clip_grad_norm(model, max_norm)
|
||||
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):
|
||||
@@ -309,3 +311,97 @@ class FSDPExecutor(BaseExecutor):
|
||||
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).
|
||||
|
||||
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.
|
||||
Original ``Parameter`` objects are preserved (as DTensors) — no
|
||||
``FlatParameter``, no ``use_orig_params=True`` hack.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
grad_accum_steps: int = 1,
|
||||
mesh: Optional[Any] = None,
|
||||
mp_policy: Optional[Any] = None,
|
||||
reshard_after_forward: bool = True,
|
||||
):
|
||||
super().__init__(grad_accum_steps=grad_accum_steps)
|
||||
self._mesh = mesh
|
||||
self._mp_policy = mp_policy
|
||||
self._reshard_after_forward = reshard_after_forward
|
||||
|
||||
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")
|
||||
return model
|
||||
|
||||
kwargs = dict(
|
||||
mesh=self._mesh,
|
||||
mp_policy=self._mp_policy,
|
||||
reshard_after_forward=self._reshard_after_forward,
|
||||
)
|
||||
kwargs = {k: v for k, v in kwargs.items() if v is not None}
|
||||
|
||||
for child in model.children():
|
||||
if isinstance(child, nn.ModuleList):
|
||||
for sub in child:
|
||||
fully_shard(sub, **kwargs)
|
||||
else:
|
||||
fully_shard(child, **kwargs)
|
||||
|
||||
logger.info(
|
||||
"FSDP2 wrapping applied to %d direct children (root skipped for ABC compat)",
|
||||
len(list(model.children())),
|
||||
)
|
||||
return model
|
||||
|
||||
@contextmanager
|
||||
def _no_sync(self, model: nn.Module):
|
||||
fsdp_modules = [m for m in model.modules() if isinstance(m, FSDPModule)]
|
||||
if fsdp_modules:
|
||||
for m in fsdp_modules:
|
||||
m.set_requires_gradient_sync(False, recurse=True)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
for m in fsdp_modules:
|
||||
m.set_requires_gradient_sync(True, recurse=True)
|
||||
else:
|
||||
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)
|
||||
|
||||
def unwrap_model(self, model: nn.Module):
|
||||
if not self.use_distributed:
|
||||
return model.state_dict()
|
||||
|
||||
if get_rank() != 0:
|
||||
return None
|
||||
|
||||
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()
|
||||
}
|
||||
|
||||
for module in model.modules():
|
||||
if isinstance(module, FSDPModule):
|
||||
module.reshard()
|
||||
|
||||
return result
|
||||
|
||||
@@ -1,5 +1,8 @@
|
||||
import logging
|
||||
import os
|
||||
import signal
|
||||
import socket
|
||||
import threading
|
||||
from abc import ABC, abstractmethod
|
||||
from contextlib import contextmanager
|
||||
from functools import wraps
|
||||
@@ -9,6 +12,10 @@ import torch
|
||||
import torch.distributed as dist
|
||||
import torch.multiprocessing as mp
|
||||
|
||||
from astrai.parallel.signal_handler import install_early_signal_handlers
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def find_free_port() -> str:
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||
@@ -115,6 +122,7 @@ def _run_single_rank(
|
||||
func: Callable,
|
||||
kwargs: dict,
|
||||
):
|
||||
install_early_signal_handlers()
|
||||
with setup_parallel(
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
@@ -155,6 +163,7 @@ class TorchrunStrategy(LaunchStrategy):
|
||||
"""External orchestrator (torchrun, SLURM, K8s) — env vars pre-set."""
|
||||
|
||||
def launch(self, func: Callable, **kwargs):
|
||||
install_early_signal_handlers()
|
||||
rank = int(os.environ["RANK"])
|
||||
world_size = int(os.environ["WORLD_SIZE"])
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", rank))
|
||||
@@ -188,6 +197,7 @@ class LocalStrategy(LaunchStrategy):
|
||||
_run_single_rank(0, *args)
|
||||
return
|
||||
|
||||
install_early_signal_handlers()
|
||||
ctx = mp.start_processes(
|
||||
_run_single_rank,
|
||||
args=args,
|
||||
@@ -195,14 +205,46 @@ class LocalStrategy(LaunchStrategy):
|
||||
start_method=self.start_method,
|
||||
join=False,
|
||||
)
|
||||
|
||||
parent_stop = threading.Event()
|
||||
original_handlers = {}
|
||||
|
||||
def _parent_handler(signum, frame):
|
||||
sig = signal.Signals(signum)
|
||||
logger.warning(
|
||||
"Parent (pid=%d) received %s, forwarding to children...",
|
||||
os.getpid(),
|
||||
sig.name,
|
||||
)
|
||||
parent_stop.set()
|
||||
for p in ctx.processes:
|
||||
if p.is_alive():
|
||||
p.terminate()
|
||||
|
||||
for sig in (signal.SIGTERM, signal.SIGINT):
|
||||
prev = signal.signal(sig, _parent_handler)
|
||||
if prev not in (signal.SIG_DFL, signal.SIG_IGN, None, _parent_handler):
|
||||
original_handlers[sig] = prev
|
||||
|
||||
try:
|
||||
while not ctx.join():
|
||||
while not ctx.join() and not parent_stop.is_set():
|
||||
pass
|
||||
except BaseException:
|
||||
logger.warning(
|
||||
"Parent received unexpected exception, terminating children..."
|
||||
)
|
||||
for p in ctx.processes:
|
||||
p.terminate()
|
||||
ctx.join()
|
||||
if p.is_alive():
|
||||
p.terminate()
|
||||
raise
|
||||
finally:
|
||||
for sig, handler in original_handlers.items():
|
||||
signal.signal(sig, handler)
|
||||
|
||||
for p in ctx.processes:
|
||||
p.join()
|
||||
|
||||
ctx.join()
|
||||
|
||||
|
||||
def _detect_launcher() -> str:
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
import logging
|
||||
import os
|
||||
import signal
|
||||
import threading
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_early_stop = threading.Event()
|
||||
_active_context = None
|
||||
|
||||
|
||||
def _early_handler(signum: int, frame):
|
||||
sig = signal.Signals(signum)
|
||||
logger.warning(
|
||||
"Received %s (pid=%d), requesting graceful training stop...",
|
||||
sig.name,
|
||||
os.getpid(),
|
||||
)
|
||||
_early_stop.set()
|
||||
if _active_context is not None:
|
||||
_active_context.request_stop()
|
||||
|
||||
|
||||
def install_early_signal_handlers():
|
||||
for sig in (signal.SIGTERM, signal.SIGINT):
|
||||
signal.signal(sig, _early_handler)
|
||||
_unblock_signals()
|
||||
|
||||
|
||||
def _unblock_signals():
|
||||
try:
|
||||
mask = signal.pthread_sigmask(signal.SIG_BLOCK, set())
|
||||
blocked = {signal.SIGTERM, signal.SIGINT} & mask
|
||||
if blocked:
|
||||
signal.pthread_sigmask(signal.SIG_UNBLOCK, blocked)
|
||||
except (AttributeError, OSError):
|
||||
pass
|
||||
|
||||
|
||||
def register_signal_handlers(context):
|
||||
global _active_context
|
||||
_active_context = context
|
||||
for sig in (signal.SIGTERM, signal.SIGINT):
|
||||
signal.signal(sig, _early_handler)
|
||||
if _early_stop.is_set():
|
||||
context.request_stop()
|
||||
logger.warning("Signal was received during initialization, stopping...")
|
||||
|
||||
|
||||
def unregister_signal_handlers():
|
||||
global _active_context
|
||||
_active_context = None
|
||||
_early_stop.clear()
|
||||
@@ -94,6 +94,97 @@ class SectionRenderer:
|
||||
|
||||
return all_ids, loss_mask
|
||||
|
||||
def process_sections_batch(
|
||||
self,
|
||||
items: list[dict],
|
||||
sections: list,
|
||||
config,
|
||||
tokenizer,
|
||||
*,
|
||||
is_top_level=False,
|
||||
filter_text=True,
|
||||
):
|
||||
"""Render and tokenize a group of records with batched Rust tokenization."""
|
||||
has_template = any(s.get("template") for s in sections)
|
||||
is_text_config = not has_template and all(
|
||||
s["action"] == "train" for s in sections
|
||||
)
|
||||
plans: list[list[tuple[str, str, bool]]] = []
|
||||
|
||||
for item in items:
|
||||
plan: list[tuple[str, str, bool]] = []
|
||||
first_section = True
|
||||
for sec in sections:
|
||||
field = sec["field"]
|
||||
action = sec["action"]
|
||||
use_template = sec.get("template", False)
|
||||
add_special = sec.get(
|
||||
"add_special_tokens", not use_template and first_section
|
||||
)
|
||||
|
||||
if use_template:
|
||||
messages = item.get(field)
|
||||
if not isinstance(messages, list) or not messages:
|
||||
continue
|
||||
for msg in messages:
|
||||
role = msg.get("role", "")
|
||||
rendered = tokenizer.apply_chat_template(
|
||||
[msg], tokenize=False, add_generation_prompt=False
|
||||
)
|
||||
plan.append(
|
||||
(rendered, _resolve_action(action, role, config), False)
|
||||
)
|
||||
else:
|
||||
text = str(item.get(field, ""))
|
||||
if not text.strip():
|
||||
continue
|
||||
if is_text_config and filter_text:
|
||||
pp = config.preprocessing
|
||||
if pp.min_chars > 0 and len(text) < pp.min_chars:
|
||||
continue
|
||||
if len(text) > pp.max_chars:
|
||||
continue
|
||||
plan.append((text, action, add_special))
|
||||
|
||||
first_section = False
|
||||
plans.append(plan)
|
||||
|
||||
encoded: dict[tuple[int, int], list[int]] = {}
|
||||
for add_special in (False, True):
|
||||
refs = [
|
||||
(item_idx, unit_idx, text)
|
||||
for item_idx, plan in enumerate(plans)
|
||||
for unit_idx, (text, _, add) in enumerate(plan)
|
||||
if add == add_special
|
||||
]
|
||||
if not refs:
|
||||
continue
|
||||
ids_batch = tokenizer.encode(
|
||||
[text for _, _, text in refs], add_special_tokens=add_special
|
||||
)
|
||||
for (item_idx, unit_idx, _), ids in zip(refs, ids_batch):
|
||||
encoded[(item_idx, unit_idx)] = ids
|
||||
|
||||
outputs = []
|
||||
max_len = config.preprocessing.max_seq_len
|
||||
for item_idx, plan in enumerate(plans):
|
||||
all_ids = []
|
||||
loss_mask = []
|
||||
if is_top_level and has_template and tokenizer.bos_token_id is not None:
|
||||
all_ids.append(tokenizer.bos_token_id)
|
||||
loss_mask.append(0)
|
||||
for unit_idx, (_, action, _) in enumerate(plan):
|
||||
ids = encoded[(item_idx, unit_idx)]
|
||||
all_ids.extend(ids)
|
||||
loss_mask.extend([1 if action == "train" else 0] * len(ids))
|
||||
all_ids = all_ids[:max_len]
|
||||
loss_mask = loss_mask[: len(all_ids)]
|
||||
if not all_ids or (is_top_level and has_template and len(all_ids) <= 1):
|
||||
outputs.append((None, None))
|
||||
else:
|
||||
outputs.append((all_ids, loss_mask))
|
||||
return outputs
|
||||
|
||||
def process_list_field(self, item: dict, sections: list, config, tokenizer):
|
||||
"""Tokenize a list-valued field, preserving per-element boundaries.
|
||||
|
||||
@@ -147,6 +238,42 @@ class SectionRenderer:
|
||||
return None, None
|
||||
return per_item_ids, per_item_masks
|
||||
|
||||
def process_list_field_batch(self, items, sections, config, tokenizer):
|
||||
per_item_ids = [[] for _ in items]
|
||||
per_item_masks = [[] for _ in items]
|
||||
|
||||
for sec in sections:
|
||||
wrappers = []
|
||||
owners = []
|
||||
field = sec["field"]
|
||||
for item_idx, item in enumerate(items):
|
||||
values = item.get(field)
|
||||
if not isinstance(values, list):
|
||||
continue
|
||||
for val in values:
|
||||
if sec.get("template", False) and not isinstance(val, list):
|
||||
continue
|
||||
wrappers.append({field: val if isinstance(val, list) else str(val)})
|
||||
owners.append(item_idx)
|
||||
|
||||
rendered = self.process_sections_batch(
|
||||
wrappers,
|
||||
[sec],
|
||||
config,
|
||||
tokenizer,
|
||||
is_top_level=False,
|
||||
filter_text=False,
|
||||
)
|
||||
for owner, (ids, mask) in zip(owners, rendered):
|
||||
if ids:
|
||||
per_item_ids[owner].append(ids)
|
||||
per_item_masks[owner].append(mask)
|
||||
|
||||
return [
|
||||
(ids, masks) if ids else (None, None)
|
||||
for ids, masks in zip(per_item_ids, per_item_masks)
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def is_value_section(sections: list) -> bool:
|
||||
return len(sections) == 1 and sections[0].get("action") == "value"
|
||||
@@ -214,6 +341,9 @@ class BaseMaskBuilder(ABC):
|
||||
@abstractmethod
|
||||
def build(self, item: dict, config, tokenizer) -> Optional[dict]: ...
|
||||
|
||||
def build_batch(self, items: list[dict], config, tokenizer) -> list[Optional[dict]]:
|
||||
return [self.build(item, config, tokenizer) for item in items]
|
||||
|
||||
|
||||
class MaskBuilderFactory(BaseFactory["BaseMaskBuilder"]):
|
||||
pass
|
||||
@@ -248,6 +378,27 @@ class SingleOutputMaskBuilder(BaseMaskBuilder):
|
||||
result["loss_mask"] = mask
|
||||
return result
|
||||
|
||||
def build_batch(self, items, config, tokenizer):
|
||||
sections = config.input.sections
|
||||
if not sections:
|
||||
return [None] * len(items)
|
||||
rendered = self.renderer.process_sections_batch(
|
||||
items, sections, config, tokenizer, is_top_level=True
|
||||
)
|
||||
results = []
|
||||
for item, (ids, mask) in zip(items, rendered):
|
||||
if ids is None:
|
||||
results.append(None)
|
||||
continue
|
||||
result = {
|
||||
"sequence": ids,
|
||||
"domain": _extract_domain(item, config.output.domain_key),
|
||||
}
|
||||
if not all(m == 1 for m in mask):
|
||||
result["loss_mask"] = mask
|
||||
results.append(result)
|
||||
return results
|
||||
|
||||
|
||||
@MaskBuilderFactory.register("multi")
|
||||
class MultiOutputMaskBuilder(BaseMaskBuilder):
|
||||
@@ -317,6 +468,49 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
|
||||
result["domain"] = _extract_domain(item, config.output.domain_key)
|
||||
return result
|
||||
|
||||
def build_batch(self, items, config, tokenizer):
|
||||
sources_spec = getattr(config.input, "sources", None)
|
||||
if not sources_spec:
|
||||
return [None] * len(items)
|
||||
|
||||
results = [{} for _ in items]
|
||||
for output_key, spec in sources_spec.items():
|
||||
sections = spec.get("sections", [])
|
||||
if not sections:
|
||||
continue
|
||||
if self.renderer.is_value_section(sections):
|
||||
for item, result in zip(items, results):
|
||||
value = self.renderer.extract_raw_value(item, sections)
|
||||
if value is not None:
|
||||
result[output_key] = value
|
||||
continue
|
||||
|
||||
mask_key = spec.get("mask_key", f"{output_key}_mask")
|
||||
if spec.get("list_field", False):
|
||||
rendered = self.renderer.process_list_field_batch(
|
||||
items, sections, config, tokenizer
|
||||
)
|
||||
else:
|
||||
rendered = self.renderer.process_sections_batch(
|
||||
items, sections, config, tokenizer, is_top_level=True
|
||||
)
|
||||
|
||||
for result, (ids, mask) in zip(results, rendered):
|
||||
if ids is None:
|
||||
continue
|
||||
result[output_key] = ids
|
||||
if spec.get("list_field", False) or not all(m == 1 for m in mask):
|
||||
result[mask_key] = mask
|
||||
elif "mask_key" in spec:
|
||||
result[mask_key] = mask
|
||||
|
||||
return [
|
||||
({**result, "domain": _extract_domain(item, config.output.domain_key)})
|
||||
if result
|
||||
else None
|
||||
for item, result in zip(items, results)
|
||||
]
|
||||
|
||||
|
||||
@MaskBuilderFactory.register("sectioned")
|
||||
class SectionedMaskBuilder(BaseMaskBuilder):
|
||||
@@ -335,3 +529,9 @@ class SectionedMaskBuilder(BaseMaskBuilder):
|
||||
if sources_spec:
|
||||
return self._multi.build(item, config, tokenizer)
|
||||
return self._single.build(item, config, tokenizer)
|
||||
|
||||
def build_batch(self, items, config, tokenizer):
|
||||
sources_spec = getattr(config.input, "sources", None)
|
||||
if sources_spec:
|
||||
return self._multi.build_batch(items, config, tokenizer)
|
||||
return self._single.build_batch(items, config, tokenizer)
|
||||
|
||||
@@ -23,7 +23,6 @@ import tqdm
|
||||
from astrai.config.preprocess_config import PipelineConfig
|
||||
from astrai.preprocessing.core import (
|
||||
build_preprocessing_components,
|
||||
iter_raw_records,
|
||||
primary_ids,
|
||||
)
|
||||
from astrai.preprocessing.packing import PackingStrategyFactory
|
||||
@@ -81,6 +80,9 @@ class Pipeline:
|
||||
def transform(self, item: dict) -> Optional[dict]:
|
||||
return self.mask_builder.build(item, self.config, self.tokenizer)
|
||||
|
||||
def transform_batch(self, items: list[dict]) -> list[Optional[dict]]:
|
||||
return self.mask_builder.build_batch(items, self.config, self.tokenizer)
|
||||
|
||||
def run(self):
|
||||
domains: dict = defaultdict(lambda: defaultdict(list))
|
||||
total_tokens = 0
|
||||
@@ -89,39 +91,55 @@ class Pipeline:
|
||||
|
||||
pp = self.config.preprocessing
|
||||
|
||||
for item in tqdm.tqdm(
|
||||
self._iter_items(), desc="Tokenizing", unit="docs", mininterval=0.5
|
||||
):
|
||||
if pp.max_items and count >= pp.max_items:
|
||||
break
|
||||
|
||||
progress = tqdm.tqdm(desc="Tokenizing", unit="docs", mininterval=0.5)
|
||||
stop = False
|
||||
for items in self._iter_batches(pp.batch_size):
|
||||
progress.update(len(items))
|
||||
try:
|
||||
result = self.transform(item)
|
||||
results = self.transform_batch(items)
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Failed to process item #%d, skipping", count + 1, exc_info=True
|
||||
"Failed to process batch, retrying records individually",
|
||||
exc_info=True,
|
||||
)
|
||||
continue
|
||||
if result is None:
|
||||
continue
|
||||
results = []
|
||||
for item in items:
|
||||
try:
|
||||
results.append(self.transform(item))
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Failed to process item, skipping", exc_info=True
|
||||
)
|
||||
results.append(None)
|
||||
|
||||
domain = result.pop("domain", "__default__")
|
||||
ids = primary_ids(result)
|
||||
if not ids:
|
||||
continue
|
||||
for result in results:
|
||||
if pp.max_items and count >= pp.max_items:
|
||||
stop = True
|
||||
break
|
||||
if result is None:
|
||||
continue
|
||||
|
||||
bucket = domains[domain]
|
||||
self._align_bucket(bucket, result, ids)
|
||||
for key, val in result.items():
|
||||
bucket[key].append(val)
|
||||
domain = result.pop("domain", "__default__")
|
||||
ids = primary_ids(result)
|
||||
if not ids:
|
||||
continue
|
||||
|
||||
count += 1
|
||||
total_tokens += len(ids)
|
||||
bucket = domains[domain]
|
||||
self._align_bucket(bucket, result, ids)
|
||||
for key, val in result.items():
|
||||
bucket[key].append(val)
|
||||
|
||||
if total_tokens >= self.config.output.max_tokens_per_shard:
|
||||
self._flush(domains, shard_idx)
|
||||
domains.clear()
|
||||
total_tokens = 0
|
||||
count += 1
|
||||
total_tokens += len(ids)
|
||||
|
||||
if total_tokens >= self.config.output.max_tokens_per_shard:
|
||||
self._flush(domains, shard_idx)
|
||||
domains.clear()
|
||||
total_tokens = 0
|
||||
if stop:
|
||||
break
|
||||
|
||||
progress.close()
|
||||
|
||||
if total_tokens > 0:
|
||||
self._flush(domains, shard_idx)
|
||||
@@ -150,6 +168,17 @@ class Pipeline:
|
||||
continue
|
||||
yield json.loads(line)
|
||||
|
||||
def _iter_batches(self, batch_size: int):
|
||||
batch_size = max(1, batch_size)
|
||||
batch = []
|
||||
for item in self._iter_items():
|
||||
batch.append(item)
|
||||
if len(batch) >= batch_size:
|
||||
yield batch
|
||||
batch = []
|
||||
if batch:
|
||||
yield batch
|
||||
|
||||
def _flush(self, domains, shard_idx):
|
||||
for domain, keys in domains.items():
|
||||
idx = shard_idx[domain]
|
||||
|
||||
@@ -100,7 +100,7 @@ def load_bin(file_path: str) -> Dict[str, List[Tensor]]:
|
||||
arr = np.memmap(
|
||||
os.path.join(file_path, f"{key}.bin"),
|
||||
dtype=info["dtype"],
|
||||
mode="r",
|
||||
mode="c",
|
||||
shape=tuple(info["shape"]),
|
||||
)
|
||||
segments[key] = [torch.from_numpy(arr)]
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
from astrai.tokenize.chat_template import ChatTemplate, MessageType
|
||||
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||
from astrai.tokenize.tokenizer import AutoTokenizer, Message, Messages
|
||||
|
||||
__all__ = [
|
||||
"AutoTokenizer",
|
||||
"ChatTemplate",
|
||||
"MessageType",
|
||||
"Message",
|
||||
"Messages",
|
||||
]
|
||||
|
||||
@@ -10,6 +10,12 @@ from tokenizers import Tokenizer
|
||||
|
||||
from astrai.tokenize.chat_template import ChatTemplate
|
||||
|
||||
Message = Dict[str, str]
|
||||
"""Single chat message with ``role`` and ``content`` keys."""
|
||||
|
||||
Messages = List[Message]
|
||||
"""Single conversation — a list of messages."""
|
||||
|
||||
|
||||
class AutoTokenizer:
|
||||
"""Base tokenizer class with automatic loading support"""
|
||||
@@ -120,7 +126,16 @@ class AutoTokenizer:
|
||||
is_pretokenized: bool = False,
|
||||
add_special_tokens: bool = True,
|
||||
) -> List:
|
||||
"""Encode text to tokens or token IDs."""
|
||||
"""Encode text to token IDs.
|
||||
|
||||
Accepts both single strings and batches:
|
||||
|
||||
- ``encode("hello")`` → ``[123, 456]``
|
||||
- ``encode(["hello", "world"])`` → ``[[123, 456], [789]]``
|
||||
|
||||
Batches are tokenised in parallel via the Rust backend's
|
||||
``encode_batch`` (uses all available CPU cores).
|
||||
"""
|
||||
if self._tokenizer is None:
|
||||
raise RuntimeError(
|
||||
"Tokenizer not initialized. Load or create a tokenizer first."
|
||||
@@ -133,15 +148,13 @@ class AutoTokenizer:
|
||||
add_special_tokens=add_special_tokens,
|
||||
)
|
||||
return encoded.ids if out_ids else encoded.tokens
|
||||
else:
|
||||
encoded_list = self._tokenizer.encode_batch(
|
||||
tokens,
|
||||
is_pretokenized=is_pretokenized,
|
||||
add_special_tokens=add_special_tokens,
|
||||
)
|
||||
return [
|
||||
encoded.ids if out_ids else encoded.tokens for encoded in encoded_list
|
||||
]
|
||||
|
||||
encoded_list = self._tokenizer.encode_batch(
|
||||
tokens,
|
||||
is_pretokenized=is_pretokenized,
|
||||
add_special_tokens=add_special_tokens,
|
||||
)
|
||||
return [encoded.ids if out_ids else encoded.tokens for encoded in encoded_list]
|
||||
|
||||
def decode(self, tokens: List[int], skip_special_tokens: bool = True) -> str:
|
||||
"""Decode token IDs to text."""
|
||||
@@ -227,45 +240,63 @@ class AutoTokenizer:
|
||||
|
||||
def apply_chat_template(
|
||||
self,
|
||||
messages: List[Dict[str, str]],
|
||||
messages: Union[Messages, List[Messages]],
|
||||
system_prompt: Optional[str] = None,
|
||||
tokenize: bool = True,
|
||||
add_generation_prompt: bool = True,
|
||||
**kwargs,
|
||||
) -> Union[str, List[int]]:
|
||||
"""
|
||||
Apply the chat template to messages and optionally tokenize the result.
|
||||
) -> Union[str, List[int], List[str], List[List[int]]]:
|
||||
"""Apply the chat template and optionally tokenize.
|
||||
|
||||
Accepts both single conversations and batches:
|
||||
|
||||
- ``apply_chat_template([msg1, msg2])`` → ``"..."`` or ``[ids]``
|
||||
- ``apply_chat_template([[msg1, msg2], [msg3]])`` → ``["..", ".."]``
|
||||
or ``[[ids], [ids]]``
|
||||
|
||||
Batches render each conversation list and tokenise all at once via
|
||||
:meth:`encode` (``List[str]`` → Rust parallel ``encode_batch``).
|
||||
|
||||
Args:
|
||||
messages: List of message dicts with 'role' and 'content'.
|
||||
system_prompt: Optional system prompt string (auto-converted to first message).
|
||||
messages: Single conversation (``Messages``) or batch of
|
||||
conversations (``BatchMessages``).
|
||||
system_prompt: Optional system prompt prepended (single mode only).
|
||||
tokenize: Whether to return token IDs (True) or raw string (False).
|
||||
add_generation_prompt: Whether to add the generation prompt (default: True).
|
||||
**kwargs: Additional variables to pass to the template.
|
||||
add_generation_prompt: Whether to add the generation prompt.
|
||||
**kwargs: Additional template variables.
|
||||
|
||||
Returns:
|
||||
Either the rendered string or list of token IDs.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If chat template is not set.
|
||||
Single mode: ``str`` or ``List[int]``.
|
||||
Batch mode: ``List[str]`` or ``List[List[int]]``.
|
||||
"""
|
||||
if self._chat_template is None:
|
||||
raise RuntimeError(
|
||||
"Chat template not set. Use set_chat_template() to set a template first."
|
||||
)
|
||||
|
||||
# Auto-convert system_prompt to first message if provided
|
||||
is_batch = bool(messages) and isinstance(messages[0], list)
|
||||
|
||||
if is_batch:
|
||||
rendered = [
|
||||
self._chat_template.render(
|
||||
messages=msgs,
|
||||
add_generation_prompt=add_generation_prompt,
|
||||
**kwargs,
|
||||
)
|
||||
for msgs in messages
|
||||
]
|
||||
if tokenize:
|
||||
return self.encode(rendered) # List[str] → batch encode
|
||||
return rendered
|
||||
|
||||
# Single conversation
|
||||
if system_prompt:
|
||||
messages = [{"role": "system", "content": system_prompt}] + list(messages)
|
||||
|
||||
# Render the template
|
||||
rendered = self._chat_template.render(
|
||||
messages=messages,
|
||||
add_generation_prompt=add_generation_prompt,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
if tokenize:
|
||||
return self.encode(rendered)
|
||||
|
||||
return rendered
|
||||
|
||||
@@ -0,0 +1,421 @@
|
||||
"""Online rollout runner for RL training.
|
||||
|
||||
Provides:
|
||||
- :class:`RawRollout` — generation output container (no reward yet)
|
||||
- :class:`RolloutResult` — a :class:`RawRollout` with rewards attached
|
||||
- :class:`BaseRewardModel` — pluggable reward interface
|
||||
- :class:`RolloutGenerator` — KV-cache-backed generation of grouped
|
||||
responses + decoding (no reward); delegates the generation loop to
|
||||
:class:`~astrai.inference.core.scheduler.InferenceScheduler.run_batch`
|
||||
so rollout and the production inference server share one code path
|
||||
- :class:`RolloutRunner` — orchestrates generation + scoring with a
|
||||
step-driven cache; its ``__call__`` returns ``(RolloutResult, is_fresh)``
|
||||
so callers do not need to rely on object identity to detect refreshes.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.inference.core.scheduler import InferenceScheduler
|
||||
|
||||
|
||||
@dataclass(kw_only=True)
|
||||
class RawRollout:
|
||||
"""Generation output before reward scoring.
|
||||
|
||||
Produced by :class:`RolloutGenerator`; consumed by :class:`RolloutRunner`
|
||||
to assemble a :class:`RolloutResult` once rewards are attached.
|
||||
|
||||
Fields are designed to cover all common RL algorithms:
|
||||
GRPO, PPO, Online DPO, Rejection Sampling, etc.
|
||||
|
||||
Fields:
|
||||
prompts: Tokenized prompts, shape ``[B, P_len]``.
|
||||
prompt_mask: Boolean mask for real prompt tokens, shape ``[B, P_len]``.
|
||||
responses: Generated response token IDs, shape ``[B, G, R_max]``.
|
||||
response_mask: Boolean mask for real (non-pad) response tokens,
|
||||
shape ``[B, G, R_max]``.
|
||||
logprobs_old: Per-token log-probs under the behaviour policy,
|
||||
shape ``[B, G, R_max]``.
|
||||
prompt_texts: Decoded prompt strings (for reward models that
|
||||
need text).
|
||||
response_texts: Decoded response strings, shape ``[B, G]``
|
||||
(for reward models).
|
||||
"""
|
||||
|
||||
prompts: Tensor
|
||||
prompt_mask: Tensor
|
||||
responses: Tensor
|
||||
response_mask: Tensor
|
||||
logprobs_old: Tensor
|
||||
prompt_texts: List[str] = field(default_factory=list)
|
||||
response_texts: List[List[str]] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass(kw_only=True)
|
||||
class RolloutResult(RawRollout):
|
||||
"""A :class:`RawRollout` with reward scoring attached.
|
||||
|
||||
Produced by :class:`RolloutRunner` once the :class:`BaseRewardModel`
|
||||
has scored the decoded responses.
|
||||
|
||||
Fields:
|
||||
rewards: Reward per response, shape ``[B, G]``.
|
||||
"""
|
||||
|
||||
rewards: Tensor
|
||||
|
||||
|
||||
class BaseRewardModel(ABC):
|
||||
"""Pluggable reward model interface.
|
||||
|
||||
Subclasses should implement ``score()`` to return a ``[B, G]`` float
|
||||
tensor of rewards. Implementations can be:
|
||||
* A loaded reward model (e.g. ArmoRM, Skywork-Reward)
|
||||
* An external API call
|
||||
* A rule-based function (format, length, keyword matching)
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def score(self, prompts: List[str], responses: List[List[str]]) -> Tensor:
|
||||
"""Score each generated response.
|
||||
|
||||
Args:
|
||||
prompts: Raw prompt strings, length ``B``.
|
||||
responses: Generated response strings, shape ``[B, G]``.
|
||||
|
||||
Returns:
|
||||
Float tensor of shape ``[B, G]``.
|
||||
"""
|
||||
...
|
||||
|
||||
|
||||
_PAD = 0
|
||||
|
||||
|
||||
class RolloutGenerator:
|
||||
"""Pure generation + decoding for a group of responses per prompt.
|
||||
|
||||
Delegates the prefill/decode loop to
|
||||
:meth:`~astrai.inference.core.scheduler.InferenceScheduler.run_batch`,
|
||||
which uses a real KV cache (no O(n²) recompute). Has no dependency
|
||||
on any reward model; can be reused in isolation for offline
|
||||
generation, qualitative sampling, or eval pipelines.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
scheduler: InferenceScheduler,
|
||||
tokenizer,
|
||||
max_tokens: int = 1024,
|
||||
group_size: int = 8,
|
||||
temperature: float = 1.0,
|
||||
top_k: int = 0,
|
||||
top_p: float = 1.0,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
):
|
||||
self.scheduler = scheduler
|
||||
self.tokenizer = tokenizer
|
||||
self.max_tokens = max_tokens
|
||||
self.group_size = group_size
|
||||
self.temperature = temperature
|
||||
self.top_k = top_k
|
||||
self.top_p = top_p
|
||||
self.frequency_penalty = frequency_penalty
|
||||
self.rep_window = rep_window
|
||||
|
||||
@torch.no_grad()
|
||||
def generate(self, batch: Dict) -> RawRollout:
|
||||
"""Expand prompts by ``group_size`` and generate one response each.
|
||||
|
||||
Accepted batch formats (per sample, repeated B times):
|
||||
|
||||
- **messages**: ``{"messages": [{"role": "user", "content": "..."}, ...]}``
|
||||
- **instruction + input + output**: ``{"instruction": "...",
|
||||
"input": "...", "output": "..."}`` — mapped to ``system`` /
|
||||
``user`` / ``assistant`` messages; ``input`` and ``output``
|
||||
are optional and skipped when empty.
|
||||
|
||||
Both are rendered through the tokenizer's chat template with
|
||||
``add_generation_prompt=True`` so rollout prompts match the
|
||||
format the policy was SFT-trained on.
|
||||
"""
|
||||
model = self.scheduler._executor.model
|
||||
was_training = model.training
|
||||
model.eval()
|
||||
try:
|
||||
return self._generate_eval(batch)
|
||||
finally:
|
||||
model.train(was_training)
|
||||
|
||||
def _generate_eval(self, batch: Dict) -> RawRollout:
|
||||
prompt_texts, flat_prompt_ids = self._prepare_prompts(batch)
|
||||
B = len(prompt_texts)
|
||||
G = self.group_size
|
||||
# Re-expand flat list to G copies per prompt for run_batch.
|
||||
expanded_prompt_ids: List[List[int]] = []
|
||||
for ids in flat_prompt_ids:
|
||||
expanded_prompt_ids.extend([list(ids)] * G)
|
||||
|
||||
results = self.scheduler.run_batch(
|
||||
expanded_prompt_ids,
|
||||
max_tokens=self.max_tokens,
|
||||
temperature=self.temperature,
|
||||
top_k=self.top_k,
|
||||
top_p=self.top_p,
|
||||
frequency_penalty=self.frequency_penalty,
|
||||
rep_window=self.rep_window,
|
||||
return_logprobs=True,
|
||||
)
|
||||
if len(results) != B * G:
|
||||
raise RuntimeError(
|
||||
f"Rollout scheduler returned {len(results)} results, expected {B * G}"
|
||||
)
|
||||
for token_ids, logprobs in results:
|
||||
if len(token_ids) != len(logprobs):
|
||||
raise RuntimeError(
|
||||
"Rollout scheduler returned misaligned token IDs and logprobs"
|
||||
)
|
||||
|
||||
# Each element is (token_ids, logprobs); pad to max length.
|
||||
max_len = 0
|
||||
for token_ids, _lp in results:
|
||||
max_len = max(max_len, len(token_ids))
|
||||
max_len = max(max_len, 1)
|
||||
|
||||
device = self.scheduler.device
|
||||
P_len = max(len(ids) for ids in flat_prompt_ids)
|
||||
prompts_tensor = torch.zeros(B, P_len, dtype=torch.long, device=device)
|
||||
prompt_mask = torch.zeros(B, P_len, dtype=torch.bool, device=device)
|
||||
for i, ids in enumerate(flat_prompt_ids):
|
||||
prompts_tensor[i, -len(ids) :] = torch.tensor(
|
||||
ids, dtype=torch.long, device=device
|
||||
)
|
||||
prompt_mask[i, -len(ids) :] = True
|
||||
|
||||
responses = torch.full((B, G, max_len), _PAD, dtype=torch.long, device=device)
|
||||
response_mask = torch.zeros((B, G, max_len), dtype=torch.bool, device=device)
|
||||
logprobs_old = torch.zeros((B, G, max_len), dtype=torch.float, device=device)
|
||||
|
||||
flat_idx = 0
|
||||
response_texts: List[List[str]] = [[] for _ in range(B)]
|
||||
for i in range(B):
|
||||
for g in range(G):
|
||||
token_ids, lps = results[flat_idx]
|
||||
flat_idx += 1
|
||||
n = len(token_ids)
|
||||
if n:
|
||||
responses[i, g, :n] = torch.tensor(
|
||||
token_ids, dtype=torch.long, device=device
|
||||
)
|
||||
response_mask[i, g, :n] = True
|
||||
logprobs_old[i, g, :n] = torch.tensor(
|
||||
lps, dtype=torch.float, device=device
|
||||
)
|
||||
response_texts[i].append(
|
||||
self.tokenizer.decode(token_ids, skip_special_tokens=True)
|
||||
)
|
||||
|
||||
return RawRollout(
|
||||
prompts=prompts_tensor,
|
||||
prompt_mask=prompt_mask,
|
||||
responses=responses,
|
||||
response_mask=response_mask,
|
||||
logprobs_old=logprobs_old,
|
||||
prompt_texts=prompt_texts,
|
||||
response_texts=response_texts,
|
||||
)
|
||||
|
||||
def _prepare_prompts(self, batch: Dict) -> Tuple[List[str], List[List[int]]]:
|
||||
"""Render batch prompts to ``(texts, token_id_lists)``.
|
||||
|
||||
Returns two parallel lists of length B (number of prompts in
|
||||
the batch). Dispatches by batch keys:
|
||||
|
||||
- ``"messages"``: treated as a pre-built message list per sample.
|
||||
- ``"instruction"`` (optionally ``"input"`` and ``"output"``): mapped
|
||||
to ``system`` / ``user`` / ``assistant`` messages respectively.
|
||||
|
||||
Both paths go through the tokenizer's chat template with
|
||||
``add_generation_prompt=True``.
|
||||
"""
|
||||
if "messages" in batch:
|
||||
messages_list = batch["messages"]
|
||||
elif "instruction" in batch:
|
||||
instructions = batch["instruction"]
|
||||
B = len(instructions)
|
||||
inputs = batch.get("input") or [""] * B
|
||||
outputs = batch.get("output") or [""] * B
|
||||
messages_list = [
|
||||
self._instruction_to_messages(i, u, o)
|
||||
for i, u, o in zip(instructions, inputs, outputs)
|
||||
]
|
||||
else:
|
||||
raise ValueError(
|
||||
"Rollout batch must contain either 'messages' or "
|
||||
"'instruction' (optionally 'input'/'output'); got keys: "
|
||||
f"{list(batch.keys())}"
|
||||
)
|
||||
|
||||
try:
|
||||
prompt_texts = self.tokenizer.apply_chat_template(
|
||||
messages_list, tokenize=False, add_generation_prompt=True
|
||||
)
|
||||
if (
|
||||
not isinstance(prompt_texts, list)
|
||||
or len(prompt_texts) != len(messages_list)
|
||||
or not all(isinstance(text, str) for text in prompt_texts)
|
||||
):
|
||||
raise TypeError("Tokenizer does not support batched chat templates")
|
||||
flat_prompt_ids = self.tokenizer.encode(prompt_texts)
|
||||
if len(flat_prompt_ids) != len(messages_list) or not all(
|
||||
isinstance(ids, list) for ids in flat_prompt_ids
|
||||
):
|
||||
raise TypeError("Tokenizer does not support batched encoding")
|
||||
except (TypeError, IndexError, KeyError):
|
||||
# Keep compatibility with lightweight tokenizer adapters that only
|
||||
# implement the single-conversation template API.
|
||||
prompt_texts = []
|
||||
flat_prompt_ids = []
|
||||
for messages in messages_list:
|
||||
text = self.tokenizer.apply_chat_template(
|
||||
messages, tokenize=False, add_generation_prompt=True
|
||||
)
|
||||
ids = self.tokenizer.apply_chat_template(
|
||||
messages, tokenize=True, add_generation_prompt=True
|
||||
)
|
||||
prompt_texts.append(text)
|
||||
flat_prompt_ids.append(list(ids))
|
||||
return prompt_texts, flat_prompt_ids
|
||||
|
||||
@staticmethod
|
||||
def _instruction_to_messages(
|
||||
instruction: str, inp: str = "", output: str = ""
|
||||
) -> List[Dict[str, str]]:
|
||||
"""Map instruction/input/output to chat messages.
|
||||
|
||||
Role mapping follows the convention used throughout the
|
||||
preprocessing pipeline: ``instruction`` → system, ``input`` →
|
||||
user, ``output`` → assistant. Empty fields are skipped so a
|
||||
bare instruction produces a ``[system]`` list and the chat
|
||||
template's ``add_generation_prompt`` adds the assistant header
|
||||
for sampling.
|
||||
"""
|
||||
messages: List[Dict[str, str]] = []
|
||||
if instruction:
|
||||
messages.append({"role": "system", "content": instruction})
|
||||
if inp:
|
||||
messages.append({"role": "user", "content": inp})
|
||||
if output:
|
||||
messages.append({"role": "assistant", "content": output})
|
||||
return messages
|
||||
|
||||
|
||||
class RolloutRunner:
|
||||
"""Produces :class:`RolloutResult` from a prompt batch.
|
||||
|
||||
Composes a :class:`RolloutGenerator` (generation + decoding) with a
|
||||
:class:`BaseRewardModel` (scoring). Maintains an internal cache so
|
||||
the same batch prompt can be replayed for multiple gradient steps.
|
||||
A new rollout is triggered every ``rollout_interval`` calls to
|
||||
:meth:`step` (or after :meth:`clear_cache`).
|
||||
|
||||
The ``__call__`` contract returns a ``(RolloutResult, is_fresh)``
|
||||
tuple — callers must use the boolean to detect a refreshed rollout
|
||||
rather than relying on object identity.
|
||||
|
||||
Usage::
|
||||
|
||||
generator = RolloutGenerator(policy, tokenizer, pipeline, ...)
|
||||
runner = RolloutRunner(generator, reward_model, rollout_interval=512)
|
||||
result, is_fresh = runner(prompt_batch)
|
||||
if is_fresh:
|
||||
... # e.g. sync behaviour policy
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
generator: RolloutGenerator,
|
||||
reward_model: BaseRewardModel,
|
||||
rollout_interval: int = 512,
|
||||
):
|
||||
self.generator = generator
|
||||
self.reward_model = reward_model
|
||||
self.rollout_interval = rollout_interval
|
||||
|
||||
self._cache: Optional[RolloutResult] = None
|
||||
self._cache_key = None
|
||||
self._steps_since_rollout: int = 0
|
||||
|
||||
def step(self):
|
||||
"""Advance the internal counter (call once per optimizer step)."""
|
||||
self._steps_since_rollout += 1
|
||||
|
||||
def clear_cache(self):
|
||||
"""Force next call to re-run rollout."""
|
||||
self._cache = None
|
||||
self._cache_key = None
|
||||
|
||||
@staticmethod
|
||||
def _batch_key(batch: Dict):
|
||||
"""Build a stable key for the prompt fields accepted by the generator."""
|
||||
|
||||
def freeze(value):
|
||||
if isinstance(value, dict):
|
||||
return tuple(sorted((key, freeze(val)) for key, val in value.items()))
|
||||
if isinstance(value, (list, tuple)):
|
||||
return tuple(freeze(item) for item in value)
|
||||
return value
|
||||
|
||||
fields = ("messages", "instruction", "input", "output")
|
||||
return tuple(
|
||||
(field, freeze(batch[field])) for field in fields if field in batch
|
||||
)
|
||||
|
||||
def _score(self, raw: RawRollout) -> RolloutResult:
|
||||
rewards = self.reward_model.score(raw.prompt_texts, raw.response_texts)
|
||||
if not isinstance(rewards, Tensor):
|
||||
rewards = torch.as_tensor(rewards, dtype=torch.float32)
|
||||
expected_shape = raw.responses.shape[:2]
|
||||
if rewards.shape != expected_shape:
|
||||
raise ValueError(
|
||||
f"Reward model returned shape {tuple(rewards.shape)}, "
|
||||
f"expected {tuple(expected_shape)}"
|
||||
)
|
||||
if not torch.isfinite(rewards).all():
|
||||
raise ValueError("Reward model returned non-finite values")
|
||||
device = raw.prompts.device
|
||||
return RolloutResult(
|
||||
prompts=raw.prompts,
|
||||
prompt_mask=raw.prompt_mask,
|
||||
responses=raw.responses,
|
||||
response_mask=raw.response_mask,
|
||||
rewards=rewards.to(device=device),
|
||||
logprobs_old=raw.logprobs_old,
|
||||
prompt_texts=raw.prompt_texts,
|
||||
response_texts=raw.response_texts,
|
||||
)
|
||||
|
||||
def __call__(self, batch: Dict[str, Tensor]) -> Tuple[RolloutResult, bool]:
|
||||
"""Return ``(cached or fresh) RolloutResult`` plus an ``is_fresh`` flag.
|
||||
|
||||
Triggers a new rollout when ``_steps_since_rollout >= rollout_interval``
|
||||
or when the cache is empty.
|
||||
"""
|
||||
cache_key = self._batch_key(batch)
|
||||
if (
|
||||
self._cache is None
|
||||
or cache_key != self._cache_key
|
||||
or self._steps_since_rollout >= self.rollout_interval
|
||||
):
|
||||
raw = self.generator.generate(batch)
|
||||
self._cache = self._score(raw)
|
||||
self._cache_key = cache_key
|
||||
self._steps_since_rollout = 0
|
||||
return self._cache, True
|
||||
return self._cache, False
|
||||
+159
-18
@@ -9,6 +9,7 @@ import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.trainer.rollout import RolloutResult
|
||||
|
||||
|
||||
def create_ref_model(
|
||||
@@ -28,9 +29,10 @@ def move_to_device(batch: Dict[str, Tensor], device: str) -> Dict[str, Tensor]:
|
||||
|
||||
|
||||
def get_logprobs(
|
||||
model: Union[nn.Module, Callable[..., Dict[str, Tensor]]],
|
||||
model: nn.Module,
|
||||
input_ids: Tensor,
|
||||
mask: Tensor,
|
||||
attn_mask: Tensor,
|
||||
loss_mask: Tensor,
|
||||
reduction: str,
|
||||
) -> Tensor:
|
||||
"""Compute token-wise log probabilities from model outputs.
|
||||
@@ -38,7 +40,8 @@ def get_logprobs(
|
||||
Args:
|
||||
model: The language model
|
||||
input_ids: Input token IDs of shape [batch_size, seq_len]
|
||||
mask: Attention mask of shape [batch_size, seq_len]
|
||||
attn_mask: Attention mask passed to the model (may include causal).
|
||||
loss_mask: Per-token mask for loss reduction.
|
||||
reduction: How to reduce over sequence dimension ("mean", "sum", "none")
|
||||
|
||||
Returns:
|
||||
@@ -51,9 +54,12 @@ def get_logprobs(
|
||||
)
|
||||
|
||||
shifted_input_ids = input_ids[:, 1:]
|
||||
shifted_mask = mask[:, 1:]
|
||||
shifted_loss_mask = loss_mask[:, 1:]
|
||||
|
||||
logits = model(input_ids[:, :-1], mask[:, :-1])["logits"]
|
||||
logits = model(
|
||||
input_ids[:, :-1],
|
||||
attn_mask[:, :, :-1, :-1] if attn_mask.dim() == 4 else attn_mask[:, :-1],
|
||||
)["logits"]
|
||||
log_probs = torch.log_softmax(logits.float(), dim=-1)
|
||||
|
||||
token_logprobs = torch.gather(
|
||||
@@ -61,13 +67,13 @@ def get_logprobs(
|
||||
).squeeze(-1)
|
||||
|
||||
if reduction == "mean":
|
||||
return (token_logprobs * shifted_mask).sum(dim=-1) / shifted_mask.sum(
|
||||
return (token_logprobs * shifted_loss_mask).sum(dim=-1) / shifted_loss_mask.sum(
|
||||
dim=-1
|
||||
).clamp(min=1.0)
|
||||
elif reduction == "sum":
|
||||
return (token_logprobs * shifted_mask).sum(dim=-1)
|
||||
return (token_logprobs * shifted_loss_mask).sum(dim=-1)
|
||||
else:
|
||||
return token_logprobs * shifted_mask
|
||||
return token_logprobs * shifted_loss_mask
|
||||
|
||||
|
||||
def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
|
||||
@@ -87,7 +93,15 @@ def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
|
||||
|
||||
|
||||
class BaseStrategy(ABC):
|
||||
"""Abstract base class for training strategies."""
|
||||
"""Abstract base class for training strategies.
|
||||
|
||||
When a :class:`~astrai.trainer.rollout.RolloutRunner` is injected via
|
||||
:meth:`set_rollout_runner`, the strategy transparently switches to
|
||||
online mode: each ``__call__`` produces a :class:`RolloutResult`,
|
||||
converts it to a training batch via :meth:`prepare_from_rollout`, and
|
||||
then computes the loss. Without a runner the strategy runs in
|
||||
offline mode and consumes the batch directly.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -99,6 +113,7 @@ class BaseStrategy(ABC):
|
||||
self.device = device
|
||||
self.executor = kwargs.pop("executor", None)
|
||||
self.extra_kwargs = kwargs
|
||||
self._rollout_runner = None
|
||||
|
||||
@abstractmethod
|
||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||
@@ -112,9 +127,53 @@ class BaseStrategy(ABC):
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def supports_online(self) -> bool:
|
||||
"""Whether this strategy can operate with a rollout runner.
|
||||
|
||||
Base implementation returns ``False``; strategies that implement
|
||||
:meth:`prepare_from_rollout` should override to return ``True``.
|
||||
"""
|
||||
return False
|
||||
|
||||
def set_rollout_runner(self, runner):
|
||||
"""Inject a :class:`RolloutRunner` to enable online rollout mode."""
|
||||
self._rollout_runner = runner
|
||||
|
||||
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
|
||||
"""Map a :class:`RolloutResult` to the batch layout expected by
|
||||
:meth:`compute_loss`.
|
||||
|
||||
Strategies that return ``True`` from :meth:`supports_online` must
|
||||
override this. Default raises :class:`NotImplementedError`.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
f"{type(self).__name__} does not support online rollout"
|
||||
)
|
||||
|
||||
def _on_rollout_refresh(self):
|
||||
"""Hook fired when a fresh rollout result is produced.
|
||||
|
||||
Override to refresh stale state (e.g. syncing the behaviour
|
||||
policy). Default is a no-op.
|
||||
"""
|
||||
pass
|
||||
|
||||
def on_optimizer_step(self):
|
||||
"""Advance online rollout state after a successful optimizer step."""
|
||||
if self._rollout_runner is not None:
|
||||
self._rollout_runner.step()
|
||||
|
||||
def __call__(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||
"""Allow calling strategy directly as a callable."""
|
||||
return self.compute_loss(batch)
|
||||
"""Run offline or online forward depending on runner injection."""
|
||||
if self._rollout_runner is None:
|
||||
return self.compute_loss(batch)
|
||||
|
||||
result, is_fresh = self._rollout_runner(batch)
|
||||
if is_fresh:
|
||||
self._on_rollout_refresh()
|
||||
|
||||
train_batch = self.prepare_from_rollout(result)
|
||||
return self.compute_loss(train_batch)
|
||||
|
||||
|
||||
class StrategyFactory(BaseFactory["BaseStrategy"]):
|
||||
@@ -238,13 +297,31 @@ class DPOStrategy(BaseStrategy):
|
||||
chosen_mask, rejected_mask = batch["chosen_mask"], batch["rejected_mask"]
|
||||
|
||||
concat_ids = torch.cat([chosen_ids, rejected_ids], dim=0)
|
||||
concat_mask = torch.cat([chosen_mask, rejected_mask], dim=0)
|
||||
concat_loss_mask = torch.cat([chosen_mask, rejected_mask], dim=0)
|
||||
|
||||
log_pi = get_logprobs(self.model, concat_ids, concat_mask, self.reduction)
|
||||
# Build full attention mask: key-padding + causal
|
||||
key_pad = concat_ids.bool()[:, None, None, :] # [B*2, 1, 1, S]
|
||||
S = key_pad.shape[-1]
|
||||
causal = torch.tril(
|
||||
torch.ones(S, S, dtype=torch.bool, device=concat_ids.device)
|
||||
)[None, None, :, :] # [1, 1, S, S]
|
||||
full_mask = key_pad & causal # [B*2, 1, S, S] — composed
|
||||
|
||||
log_pi = get_logprobs(
|
||||
self.model,
|
||||
concat_ids,
|
||||
full_mask,
|
||||
concat_loss_mask,
|
||||
self.reduction,
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
log_ref = get_logprobs(
|
||||
self.ref_model, concat_ids, concat_mask, self.reduction
|
||||
self.ref_model,
|
||||
concat_ids,
|
||||
full_mask,
|
||||
concat_loss_mask,
|
||||
self.reduction,
|
||||
)
|
||||
|
||||
log_pi_chosen = log_pi[: chosen_ids.shape[0]]
|
||||
@@ -260,6 +337,29 @@ class DPOStrategy(BaseStrategy):
|
||||
|
||||
return dpo_loss
|
||||
|
||||
def supports_online(self) -> bool:
|
||||
return True
|
||||
|
||||
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
|
||||
"""Pick best/worst response per prompt by reward as chosen/rejected."""
|
||||
rewards = result.rewards
|
||||
responses = result.responses
|
||||
masks = result.response_mask
|
||||
best = rewards.argmax(dim=-1)
|
||||
worst = rewards.argmin(dim=-1)
|
||||
B = responses.shape[0]
|
||||
idx = torch.arange(B, device=responses.device)
|
||||
chosen = responses[idx, best]
|
||||
chosen_mask = masks[idx, best].float()
|
||||
rejected = responses[idx, worst]
|
||||
rejected_mask = masks[idx, worst].float()
|
||||
return {
|
||||
"chosen": chosen,
|
||||
"chosen_mask": chosen_mask,
|
||||
"rejected": rejected,
|
||||
"rejected_mask": rejected_mask,
|
||||
}
|
||||
|
||||
|
||||
@StrategyFactory.register("grpo")
|
||||
class GRPOStrategy(BaseStrategy):
|
||||
@@ -314,6 +414,12 @@ class GRPOStrategy(BaseStrategy):
|
||||
responses_flat = responses.view(-1, response_len)
|
||||
masks_flat = masks.view(-1, response_len)
|
||||
prompt_expanded = prompts.unsqueeze(1).repeat(1, group_size, 1).flatten(0, 1)
|
||||
prompt_mask = batch.get("prompt_mask")
|
||||
if prompt_mask is None:
|
||||
prompt_mask = prompts.ne(0)
|
||||
prompt_mask_expanded = (
|
||||
prompt_mask.unsqueeze(1).expand(-1, group_size, -1).flatten(0, 1)
|
||||
)
|
||||
prompt_len = prompt_expanded.size(1)
|
||||
|
||||
full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-1)
|
||||
@@ -321,20 +427,32 @@ class GRPOStrategy(BaseStrategy):
|
||||
# response tokens. get_logprobs shifts the mask by one position, so
|
||||
# the first response token's logprob (predicted from the last prompt
|
||||
# token) is correctly included.
|
||||
full_masks = torch.cat([torch.zeros_like(prompt_expanded), masks_flat], dim=-1)
|
||||
full_masks = torch.cat(
|
||||
[torch.zeros_like(prompt_expanded, dtype=torch.bool), masks_flat], dim=-1
|
||||
)
|
||||
|
||||
# Build full attention mask: key-padding + causal
|
||||
key_pad = torch.cat([prompt_mask_expanded, masks_flat.bool()], dim=-1)[
|
||||
:, None, None, :
|
||||
]
|
||||
S = key_pad.shape[-1]
|
||||
causal = torch.tril(
|
||||
torch.ones(S, S, dtype=torch.bool, device=full_sequences.device)
|
||||
)[None, None, :, :]
|
||||
attn_mask = key_pad & causal
|
||||
|
||||
# get_logprobs returns [B*G, S-1] (S = prompt_len + response_len).
|
||||
# Response token logprobs occupy the last ``response_len`` positions
|
||||
# (the first response token is predicted from the last prompt token).
|
||||
token_log_probs_policy = get_logprobs(
|
||||
self.model, full_sequences, full_masks, "none"
|
||||
self.model, full_sequences, attn_mask, full_masks, "none"
|
||||
)[:, prompt_len - 1 :]
|
||||
with torch.no_grad():
|
||||
token_log_probs_old = get_logprobs(
|
||||
self.old_model, full_sequences, full_masks, "none"
|
||||
self.old_model, full_sequences, attn_mask, full_masks, "none"
|
||||
)[:, prompt_len - 1 :]
|
||||
token_log_probs_ref = get_logprobs(
|
||||
self.ref_model, full_sequences, full_masks, "none"
|
||||
self.ref_model, full_sequences, attn_mask, full_masks, "none"
|
||||
)[:, prompt_len - 1 :]
|
||||
|
||||
# Reshape to [B, G, response_len]
|
||||
@@ -371,3 +489,26 @@ class GRPOStrategy(BaseStrategy):
|
||||
total_loss = policy_loss + kl_penalty
|
||||
|
||||
return total_loss
|
||||
|
||||
def supports_online(self) -> bool:
|
||||
return True
|
||||
|
||||
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
|
||||
return {
|
||||
"prompts": result.prompts,
|
||||
"prompt_mask": result.prompt_mask,
|
||||
"responses": result.responses,
|
||||
"masks": result.response_mask,
|
||||
"rewards": result.rewards,
|
||||
}
|
||||
|
||||
def _on_rollout_refresh(self):
|
||||
"""Sync the behaviour policy whenever a fresh rollout arrives."""
|
||||
self.sync_old_model()
|
||||
|
||||
|
||||
# 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
|
||||
|
||||
+114
-52
@@ -1,3 +1,4 @@
|
||||
import threading
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Optional, Self
|
||||
@@ -8,11 +9,14 @@ from torch.utils.data import DataLoader, random_split
|
||||
|
||||
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.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
|
||||
|
||||
|
||||
@@ -27,7 +31,6 @@ class TrainContext:
|
||||
config: TrainConfig = field(default=None)
|
||||
model_config: dict = field(default_factory=dict)
|
||||
executor: BaseExecutor = field(default=None)
|
||||
|
||||
epoch: int = field(default=0)
|
||||
consumed_samples: int = field(default=0)
|
||||
loss: float = field(default=0.0)
|
||||
@@ -39,6 +42,15 @@ class TrainContext:
|
||||
rank: int = field(default=0)
|
||||
kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
_stop_event: threading.Event = field(default_factory=threading.Event)
|
||||
|
||||
@property
|
||||
def stop_requested(self) -> bool:
|
||||
return self._stop_event.is_set()
|
||||
|
||||
def request_stop(self) -> None:
|
||||
self._stop_event.set()
|
||||
|
||||
@property
|
||||
def optimizer_step(self) -> int:
|
||||
return self.consumed_samples // (
|
||||
@@ -72,61 +84,70 @@ class TrainContextBuilder:
|
||||
**cfg.executor_kwargs,
|
||||
)
|
||||
|
||||
model = cfg.model_fn()
|
||||
model = model.to(device=device)
|
||||
|
||||
model_config = {}
|
||||
if self._param_path:
|
||||
config_path = Path(self._param_path) / "config.json"
|
||||
if config_path.exists():
|
||||
model_config = load_json(config_path)
|
||||
|
||||
if not model_config and hasattr(model, "config"):
|
||||
model_config = model.config.to_dict()
|
||||
preloaded_state_dict = None
|
||||
preloaded_epoch = cfg.start_epoch
|
||||
preloaded_consumed = cfg.start_samples * get_world_size()
|
||||
preloaded_checkpoint = None
|
||||
if self._param_path:
|
||||
checkpoint = Checkpoint.load_any(self._param_path)
|
||||
if checkpoint is not None:
|
||||
preloaded_state_dict = checkpoint.state_dict
|
||||
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_checkpoint = checkpoint
|
||||
|
||||
if not model_config and hasattr(cfg.model_fn(), "config"):
|
||||
model_config = cfg.model_fn().config.to_dict()
|
||||
|
||||
def _before_wrap(m):
|
||||
m = m.to(device=device)
|
||||
if cfg.lora is not None:
|
||||
inject_lora(
|
||||
m,
|
||||
r=cfg.lora.r,
|
||||
alpha=cfg.lora.alpha,
|
||||
target_modules=set(cfg.lora.target_modules),
|
||||
)
|
||||
if preloaded_state_dict is not None:
|
||||
m.load_state_dict(preloaded_state_dict, strict=False)
|
||||
return m
|
||||
|
||||
context = TrainContext(
|
||||
model=model,
|
||||
world_size=get_world_size(),
|
||||
rank=get_rank(),
|
||||
config=cfg,
|
||||
model_config=model_config,
|
||||
executor=executor,
|
||||
epoch=preloaded_epoch,
|
||||
consumed_samples=preloaded_consumed,
|
||||
checkpoint=preloaded_checkpoint,
|
||||
)
|
||||
|
||||
if self._param_path:
|
||||
checkpoint = Checkpoint.load_any(self._param_path)
|
||||
if checkpoint is not None:
|
||||
model.load_state_dict(checkpoint.state_dict, strict=False)
|
||||
if checkpoint.config:
|
||||
context.model_config = checkpoint.config
|
||||
|
||||
if self._resume:
|
||||
context.epoch = checkpoint.epoch or cfg.start_epoch
|
||||
if checkpoint.consumed_samples > 0:
|
||||
per_step = (
|
||||
cfg.batch_per_device
|
||||
* context.world_size
|
||||
* cfg.grad_accum_steps
|
||||
)
|
||||
context.consumed_samples = (
|
||||
checkpoint.consumed_samples // per_step
|
||||
) * per_step
|
||||
else:
|
||||
context.consumed_samples = (
|
||||
cfg.start_samples * context.world_size
|
||||
)
|
||||
context.checkpoint = checkpoint
|
||||
|
||||
if cfg.lora is not None:
|
||||
inject_lora(
|
||||
model,
|
||||
r=cfg.lora.r,
|
||||
alpha=cfg.lora.alpha,
|
||||
target_modules=set(cfg.lora.target_modules),
|
||||
)
|
||||
|
||||
context.optimizer = cfg.optimizer_fn(model)
|
||||
context.scheduler = cfg.scheduler_fn(context.optimizer)
|
||||
context.model, context.optimizer, context.scheduler = executor.prepare(
|
||||
cfg.model_fn,
|
||||
cfg.optimizer_fn,
|
||||
cfg.scheduler_fn,
|
||||
before_wrap=_before_wrap,
|
||||
)
|
||||
|
||||
train_dataset = cfg.dataset
|
||||
val_dataset = cfg.val_dataset
|
||||
@@ -175,15 +196,6 @@ class TrainContextBuilder:
|
||||
collate_fn=cfg.collate_fn,
|
||||
)
|
||||
|
||||
context.model, context.optimizer, context.dataloader, context.scheduler = (
|
||||
executor.prepare(
|
||||
model,
|
||||
context.optimizer,
|
||||
context.dataloader,
|
||||
context.scheduler,
|
||||
)
|
||||
)
|
||||
|
||||
if context.checkpoint and context.checkpoint.extra:
|
||||
extra = context.checkpoint.extra
|
||||
for name in ("optimizer", "scheduler"):
|
||||
@@ -194,13 +206,22 @@ class TrainContextBuilder:
|
||||
|
||||
strategy_kwargs = dict(cfg.extra_kwargs)
|
||||
|
||||
if cfg.strategy in ("dpo", "grpo"):
|
||||
needs_ref = cfg.strategy in (
|
||||
"dpo",
|
||||
"grpo",
|
||||
"online_grpo",
|
||||
"online_dpo",
|
||||
)
|
||||
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
|
||||
|
||||
if cfg.strategy == "grpo":
|
||||
old_model = None
|
||||
if needs_old:
|
||||
old_model = create_ref_model(
|
||||
cfg.model_fn, executor.unwrap_model(context.model)
|
||||
).to(device=device)
|
||||
@@ -214,4 +235,45 @@ class TrainContextBuilder:
|
||||
**strategy_kwargs,
|
||||
)
|
||||
|
||||
# Enable online rollout when the train_type is an ``online_*`` variant.
|
||||
is_online = cfg.strategy.startswith("online_")
|
||||
if is_online:
|
||||
if not context.strategy.supports_online():
|
||||
raise ValueError(
|
||||
f"Strategy '{cfg.strategy}' does not support online rollout"
|
||||
)
|
||||
if cfg.reward_model_fn is None:
|
||||
raise ValueError("reward_model_fn is required for online RL strategies")
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(self._param_path)
|
||||
reward_model = cfg.reward_model_fn()
|
||||
|
||||
group_size = strategy_kwargs.get("group_size", 1)
|
||||
rollout_batch_size = group_size * max(1, cfg.batch_per_device)
|
||||
max_seq_len = getattr(context.model.config, "max_position_embeddings", None)
|
||||
|
||||
scheduler = InferenceScheduler(
|
||||
model=context.model,
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=rollout_batch_size,
|
||||
max_seq_len=max_seq_len,
|
||||
max_prompt_len=max_seq_len or 4096,
|
||||
)
|
||||
|
||||
generator = RolloutGenerator(
|
||||
scheduler=scheduler,
|
||||
tokenizer=tokenizer,
|
||||
max_tokens=cfg.rollout_max_tokens,
|
||||
group_size=group_size,
|
||||
temperature=cfg.rollout_temperature,
|
||||
top_k=cfg.rollout_top_k,
|
||||
top_p=cfg.rollout_top_p,
|
||||
)
|
||||
runner = RolloutRunner(
|
||||
generator=generator,
|
||||
reward_model=reward_model,
|
||||
rollout_interval=cfg.rollout_interval,
|
||||
)
|
||||
context.strategy.set_rollout_runner(runner)
|
||||
|
||||
return context
|
||||
|
||||
@@ -1,8 +1,14 @@
|
||||
import logging
|
||||
from typing import List, Optional
|
||||
|
||||
import torch.distributed as dist
|
||||
|
||||
from astrai.config import TrainConfig
|
||||
from astrai.parallel.setup import spawn_parallel_fn
|
||||
from astrai.parallel.signal_handler import (
|
||||
register_signal_handlers,
|
||||
unregister_signal_handlers,
|
||||
)
|
||||
from astrai.trainer.train_callback import (
|
||||
CallbackFactory,
|
||||
TrainCallback,
|
||||
@@ -58,6 +64,7 @@ class Trainer:
|
||||
.with_param_path(param_path, resume=resume)
|
||||
.build()
|
||||
)
|
||||
register_signal_handlers(context)
|
||||
executor = context.executor
|
||||
self._call_callbacks("on_train_begin", context)
|
||||
|
||||
@@ -65,10 +72,14 @@ class Trainer:
|
||||
context.model.train()
|
||||
|
||||
for epoch in range(context.epoch, context.config.n_epoch):
|
||||
if context.stop_requested:
|
||||
break
|
||||
context.epoch = epoch
|
||||
self._call_callbacks("on_epoch_begin", context)
|
||||
|
||||
for batch in context.dataloader:
|
||||
if context.stop_requested:
|
||||
break
|
||||
with executor.accumulate(context.model):
|
||||
self._call_callbacks("on_batch_begin", context)
|
||||
loss = context.strategy(batch)
|
||||
@@ -83,6 +94,7 @@ class Trainer:
|
||||
if executor.sync_gradients:
|
||||
self._call_callbacks("on_optimizer_step", context)
|
||||
context.optimizer.step()
|
||||
context.strategy.on_optimizer_step()
|
||||
context.optimizer.zero_grad()
|
||||
|
||||
if context.scheduler:
|
||||
@@ -90,12 +102,21 @@ class Trainer:
|
||||
|
||||
self._call_callbacks("on_epoch_end", context)
|
||||
|
||||
if context.stop_requested:
|
||||
logger.warning(
|
||||
"Training interrupted by signal, saving emergency checkpoint..."
|
||||
)
|
||||
self._call_callbacks("on_error", context)
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Training failed: %s", str(e), exc_info=True)
|
||||
self._call_callbacks("on_error", context)
|
||||
raise
|
||||
finally:
|
||||
self._call_callbacks("on_train_end", context)
|
||||
if executor.use_distributed and dist.is_initialized():
|
||||
dist.barrier()
|
||||
unregister_signal_handlers()
|
||||
|
||||
def train(self, param_path: Optional[str] = None, resume: bool = False):
|
||||
cfg = self.train_config
|
||||
|
||||
@@ -1,51 +1,6 @@
|
||||
#include "attn_decode_split_kv.cuh"
|
||||
#include "attn_dispatchers.cuh"
|
||||
#include "attn_entry_utils.cuh"
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
#include "attn_decode_split_kv_mma.cuh"
|
||||
#endif
|
||||
|
||||
// Scalar fallback: one warp per query head, split-KV across grid.z.
|
||||
static void launch_scalar_decode(AttentionParams<bf16>& p) {
|
||||
int group_size = p.q_head / p.kv_head;
|
||||
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||
alloc_split_partials(p);
|
||||
|
||||
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
||||
attn_decode_split_kv_kernel<<<dim3(p.batch * p.kv_head, 1, p.num_splits), dim3(32, group_size), smem>>>(p);
|
||||
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
// MMA head-packing requires G <= 16 (BR=16 rows). sm_80+ tensor-core
|
||||
// + cp.async wins even at G=1 (decode is memory-bound, not compute-bound).
|
||||
// STAGES=2 (double-buffer) for D<=128 (smem 16 KB); STAGES=1 for D=256
|
||||
// (double-buffer would be 32 KB, near the 48 KB static cap — keep single
|
||||
// to preserve occupancy).
|
||||
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
|
||||
static void launch_mma_decode(AttentionParams<bf16>& p) {
|
||||
int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
||||
alloc_split_partials(p);
|
||||
|
||||
attn_decode_split_kv_mma_kernel<HEAD_DIM, BC, STAGES><<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
|
||||
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
#endif
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static void dispatch_decode(AttentionParams<bf16>& p) {
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
int G = p.q_head / p.kv_head;
|
||||
if (G >= 1 && G <= 16) {
|
||||
launch_mma_decode<HEAD_DIM, 32>(p);
|
||||
return;
|
||||
}
|
||||
#endif
|
||||
launch_scalar_decode(p);
|
||||
}
|
||||
|
||||
torch::Tensor attn_decode(
|
||||
torch::Tensor q,
|
||||
torch::Tensor k,
|
||||
@@ -60,11 +15,11 @@ torch::Tensor attn_decode(
|
||||
TORCH_CHECK(p.q_len == 1, "Q seq_len must be 1");
|
||||
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
|
||||
|
||||
// O matches Q's original layout
|
||||
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
||||
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
|
||||
p.o = (bf16*)O_view.data_ptr();
|
||||
|
||||
alloc_split_partials(p);
|
||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p);
|
||||
return O;
|
||||
}
|
||||
|
||||
@@ -2,16 +2,10 @@
|
||||
#include <cuda_bf16.h>
|
||||
#include <float.h>
|
||||
#include "attn_common.h"
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
#include "attn_warp_utils.cuh"
|
||||
constexpr int DC_CHUNK = 64;
|
||||
|
||||
__device__ inline float warp_reduce_sum(float val) {
|
||||
for (int offset = 16; offset > 0; offset >>= 1)
|
||||
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
|
||||
return val;
|
||||
}
|
||||
|
||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||
__global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
||||
int batch = blockIdx.x / p.kv_head;
|
||||
int kv_head = blockIdx.x % p.kv_head;
|
||||
@@ -48,7 +42,8 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
||||
|
||||
// Load K into shared memory (gather from strided global)
|
||||
int total = this_chunk * p.head_dim;
|
||||
for (int i = threadIdx.y * 32 + lane; i < total; i += blockDim.x * blockDim.y) {
|
||||
for (int i = threadIdx.y * 32 + lane; i < total;
|
||||
i += blockDim.x * blockDim.y) {
|
||||
int s = i / p.head_dim;
|
||||
int d_dim = i % p.head_dim;
|
||||
int kv_idx = chunk_start + s;
|
||||
@@ -60,24 +55,30 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
||||
for (int s = 0; s < this_chunk; s++) {
|
||||
float partial = 0.0f;
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
partial += q_reg[i] * __bfloat162float(k_smem[s * p.head_dim + lane * hd_per_thread + i]);
|
||||
partial += q_reg[i] * __bfloat162float(
|
||||
k_smem[s * p.head_dim + lane * hd_per_thread + i]);
|
||||
partial = warp_reduce_sum(partial) * p.scale;
|
||||
|
||||
int kv_idx = chunk_start + s;
|
||||
if (p.use_mask && p.mask && !p.mask[mask_base + kv_idx])
|
||||
partial = -FLT_MAX;
|
||||
if (p.causal_offset >= 0 && kv_idx > p.causal_offset)
|
||||
partial = -FLT_MAX;
|
||||
if constexpr (HasMask) {
|
||||
if (!p.mask[mask_base + kv_idx])
|
||||
partial = -FLT_MAX;
|
||||
}
|
||||
if constexpr (IsCausal) {
|
||||
if (kv_idx > p.causal_offset)
|
||||
partial = -FLT_MAX;
|
||||
}
|
||||
|
||||
float new_m = fmaxf(m, partial);
|
||||
float alpha = expf(m - new_m);
|
||||
float beta = expf(partial - new_m);
|
||||
d = d * alpha + beta;
|
||||
|
||||
// V: stride-based read
|
||||
int v_off = kv_base + kv_idx * p.kv_stride_l + lane * hd_per_thread * p.kv_stride_d;
|
||||
int v_off = kv_base + kv_idx * p.kv_stride_l
|
||||
+ lane * hd_per_thread * p.kv_stride_d;
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
acc_reg[i] = acc_reg[i] * alpha + __bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta;
|
||||
acc_reg[i] = fmaf(acc_reg[i], alpha,
|
||||
__bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta);
|
||||
m = new_m;
|
||||
}
|
||||
__syncthreads();
|
||||
@@ -85,7 +86,7 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
||||
|
||||
// ---- write UN-normalised partials for this split ----
|
||||
size_t bh = (size_t)batch * p.q_head + q_head;
|
||||
size_t slot = bh * p.num_splits + split;
|
||||
size_t slot = bh * MAX_SPLITS + split;
|
||||
int d0 = lane * hd_per_thread;
|
||||
for (int i = 0; i < hd_per_thread; i++) {
|
||||
int dd = d0 + i;
|
||||
@@ -97,9 +98,6 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
||||
}
|
||||
}
|
||||
|
||||
// Reduce split-K partials into the final bf16 output. One block per (batch,
|
||||
// q_head); each thread folds across all splits with a single-pass
|
||||
// online-rescale reduction (expf + FMA counts halved vs 3-pass original).
|
||||
__global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
|
||||
int bh = blockIdx.x;
|
||||
int d = threadIdx.x;
|
||||
@@ -108,7 +106,7 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
|
||||
int batch = bh / p.q_head;
|
||||
int q_head = bh % p.q_head;
|
||||
|
||||
size_t split_base = (size_t)bh * p.num_splits;
|
||||
size_t split_base = (size_t)bh * MAX_SPLITS;
|
||||
const float* mlp = p.ml_part + split_base * 2;
|
||||
const float* op = p.o_part + split_base * p.head_dim;
|
||||
|
||||
@@ -118,15 +116,14 @@ __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);
|
||||
acc = acc * corr + op[s * p.head_dim + d] * e;
|
||||
l = l * corr + li * e;
|
||||
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;
|
||||
}
|
||||
|
||||
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
||||
// Stride-based output write (q_len=1 for decode, so stride_l not needed)
|
||||
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h + d * p.q_stride_d;
|
||||
p.o[o_off] = __float2bfloat16(acc * inv);
|
||||
}
|
||||
|
||||
@@ -3,85 +3,72 @@
|
||||
#include <cuda_bf16.h>
|
||||
#include "attn_common.h"
|
||||
#include "attn_mma_utils.cuh"
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
#include "attn_warp_utils.cuh"
|
||||
|
||||
// Split-K (FlashDecoding) tensor-core decode via GQA head-packing.
|
||||
// Decode has q_len == 1, so we pack G = q_head/kv_head query heads into the
|
||||
// M=16 rows of mma.sync.m16n8k16, turning G independent GEMVs into a single
|
||||
// GEMM that reuses each loaded K/V tile across all G heads.
|
||||
//
|
||||
// Decode has q_len == 1, so S = q @ K^T is a GEMV per head — no tensor-core
|
||||
// work on its own. But GQA gives us G = q_head / kv_head query heads that all
|
||||
// share one kv_head. We pack those G heads into the M=16 rows of
|
||||
// mma.sync.m16n8k16, turning G independent GEMVs into a single GEMM that
|
||||
// reuses each loaded K/V tile across all G heads (K/V load is the decode
|
||||
// bottleneck, so the reuse is the win, not the flops). The KV sequence is
|
||||
// partitioned across gridDim.z blocks so that a decode with only
|
||||
// batch*kv_head independent tasks can fill all SMs. Each (batch, kv_head,
|
||||
// split) block computes an UN-normalised partial (Oacc, m, l) over its KV
|
||||
// slice; the combine kernel below reduces across splits. Fixes the "grid too
|
||||
// small" bottleneck (0.04 waves/SM → many blocks) for long-context,
|
||||
// small-batch decode.
|
||||
|
||||
template <int HEAD_DIM, int BC, int STAGES = 2>
|
||||
// IsCausal and HasMask are compile-time bools — no runtime branch in the
|
||||
// inner compute loop.
|
||||
//
|
||||
// Traits = KernelTraits<HEAD_DIM, BC=32, WARPS=1, STAGES=<2 or 1>>.
|
||||
template <typename Traits, bool IsCausal, bool HasMask>
|
||||
__global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
||||
constexpr int KD = HEAD_DIM / 16;
|
||||
constexpr int NC8 = BC / 8;
|
||||
constexpr int KT2 = BC / 16;
|
||||
constexpr int DN8 = HEAD_DIM / 8;
|
||||
constexpr int LD = HEAD_DIM;
|
||||
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
|
||||
constexpr int VEC = 8;
|
||||
constexpr int TOTAL = BC * HEAD_DIM;
|
||||
|
||||
const int lane = threadIdx.x;
|
||||
const int gid = lane >> 2;
|
||||
const int tid4 = lane & 3;
|
||||
|
||||
const int kv_head = blockIdx.x;
|
||||
const int pass = blockIdx.x / p.kv_head;
|
||||
const int kv_head = blockIdx.x % p.kv_head;
|
||||
const int batch = blockIdx.y;
|
||||
const int split = blockIdx.z;
|
||||
const int G = p.q_head / p.kv_head;
|
||||
const int q_head0 = kv_head * G;
|
||||
|
||||
// Double-buffered shared memory for K/V (no sQ needed — Q goes direct
|
||||
// from global to registers).
|
||||
__shared__ __align__(16) bf16 sK[STAGES * BC * LD];
|
||||
__shared__ __align__(16) bf16 sV[STAGES * BC * LD];
|
||||
constexpr int MAX_G = 16;
|
||||
const int G_total = p.q_head / p.kv_head;
|
||||
const int g_begin = pass * MAX_G;
|
||||
const int G = min(MAX_G, G_total - g_begin);
|
||||
const int q_head0 = kv_head * G_total + g_begin;
|
||||
|
||||
// ---- Load Q directly from global into mma A-operand registers ----
|
||||
// Double-buffered shared memory for K/V (no sQ needed)
|
||||
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
|
||||
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
|
||||
|
||||
// Load Q directly from global into mma A-operand registers.
|
||||
// stride_row = p.q_stride_h for decode (q_len=1).
|
||||
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
|
||||
const int qra = gid;
|
||||
const int qrb = gid + 8;
|
||||
const bool va = qra < G, vb = qrb < G;
|
||||
unsigned Qa[KD][4];
|
||||
load_q_mma_frags<KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
|
||||
qra, qrb, va, vb, tid4, Qa);
|
||||
unsigned Qa[Traits::KD][4];
|
||||
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
|
||||
qra, qrb, va, vb, tid4, Qa);
|
||||
|
||||
float Oacc[DN8][4];
|
||||
#pragma unroll
|
||||
for (int j = 0; j < DN8; j++)
|
||||
float Oacc[Traits::DN8][4];
|
||||
#pragma unroll
|
||||
for (int j = 0; j < Traits::DN8; j++)
|
||||
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
||||
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
||||
|
||||
// KV: stride-based base — [batch, kv_head, kv_len, head_dim]
|
||||
const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||
const int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||
const int tiles_total = (p.kv_len + Traits::BC - 1) / Traits::BC;
|
||||
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
|
||||
const int ti_begin = split * tiles_per_split;
|
||||
const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
|
||||
const int has_mask = p.use_mask && p.mask;
|
||||
|
||||
// ---- Load tile lambda: predicated cp.async, unified full/partial ----
|
||||
// ---- Load tile lambda: predicated cp.async ----
|
||||
auto load_tile = [&](int ti, int buf) {
|
||||
int kv0 = ti * BC;
|
||||
bf16* dK = sK + buf * BC * LD;
|
||||
bf16* dV = sV + buf * BC * LD;
|
||||
#pragma unroll
|
||||
for (int i = lane * VEC; i < TOTAL; i += 32 * VEC) {
|
||||
int r = i / HEAD_DIM, d = i % HEAD_DIM;
|
||||
int kv0 = ti * Traits::BC;
|
||||
bf16* dK = sK + buf * Traits::BC * Traits::LD;
|
||||
bf16* dV = sV + buf * Traits::BC * Traits::LD;
|
||||
#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;
|
||||
int off = r * LD + swiz_col(d, r, SWIZ_MASK);
|
||||
// KV stride-based: contiguous within head_dim (stride_d == 1 typically)
|
||||
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
|
||||
int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d;
|
||||
cp_async_16_pred(&dK[off], &p.k[g_off], valid);
|
||||
cp_async_16_pred(&dV[off], &p.v[g_off], valid);
|
||||
@@ -89,50 +76,48 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
||||
cp_async_commit();
|
||||
};
|
||||
|
||||
// ---- Prologue: issue first tile load ----
|
||||
constexpr int BUF_MASK = (Traits::STAGES > 1) ? (Traits::STAGES - 1) : 0;
|
||||
|
||||
// Prologue
|
||||
if (ti_begin < ti_end) {
|
||||
load_tile(ti_begin, 0);
|
||||
}
|
||||
|
||||
for (int ti = ti_begin; ti < ti_end; ti++) {
|
||||
constexpr int BUF_MASK = (STAGES > 1) ? (STAGES - 1) : 0;
|
||||
int buf = (ti - ti_begin) & BUF_MASK;
|
||||
|
||||
// Wait for current tile, then issue next tile's prefetch (overlaps
|
||||
// with this tile's compute). Single syncwarp covers both hazards.
|
||||
// When STAGES==1, no prefetch — load happens at end of prior iter.
|
||||
cp_async_wait_group<0>();
|
||||
__syncwarp();
|
||||
if constexpr (STAGES > 1) {
|
||||
if constexpr (Traits::STAGES > 1) {
|
||||
if (ti + 1 < ti_end)
|
||||
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
|
||||
}
|
||||
|
||||
const bf16* bK = sK + buf * BC * LD;
|
||||
const bf16* bV = sV + buf * BC * LD;
|
||||
int kv0 = ti * BC;
|
||||
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
|
||||
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
|
||||
int kv0 = ti * Traits::BC;
|
||||
|
||||
float Sacc[NC8][4];
|
||||
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
|
||||
float Sacc[Traits::NC8][4];
|
||||
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
|
||||
|
||||
#pragma unroll
|
||||
for (int n8 = 0; n8 < NC8; n8++)
|
||||
for (int n8 = 0; n8 < Traits::NC8; n8++)
|
||||
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||
|
||||
// Decode: q_len=1, so qrow0=qrow1=0, mask_q_stride irrelevant
|
||||
int maxc = (p.causal_offset >= 0) ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
||||
mma_softmax_tile<NC8, DN8>(kv0, maxc, maxc,
|
||||
0, 0,
|
||||
p.mask_b_stride, 0,
|
||||
batch,
|
||||
p.mask, has_mask,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
// Decode: q_len=1, so qrow0=qrow1=0
|
||||
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,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
|
||||
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc);
|
||||
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
||||
__syncwarp();
|
||||
|
||||
if constexpr (STAGES == 1) {
|
||||
if constexpr (Traits::STAGES == 1) {
|
||||
if (ti + 1 < ti_end)
|
||||
load_tile(ti + 1, 0);
|
||||
}
|
||||
@@ -141,21 +126,21 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
||||
// ---- write UN-normalised partials for this split ----
|
||||
auto split_slot = [&](int h) -> size_t {
|
||||
size_t bh = (size_t)batch * p.q_head + h;
|
||||
return bh * p.num_splits + split;
|
||||
return bh * MAX_SPLITS + split;
|
||||
};
|
||||
#pragma unroll
|
||||
for (int dn8 = 0; dn8 < DN8; dn8++) {
|
||||
#pragma unroll
|
||||
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
|
||||
int d = dn8 * 8 + 2 * tid4;
|
||||
int r0 = gid, r1 = gid + 8;
|
||||
if (r0 < G) {
|
||||
int h = q_head0 + r0;
|
||||
float* op = p.o_part + split_slot(h) * HEAD_DIM;
|
||||
float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
|
||||
op[d] = Oacc[dn8][0];
|
||||
op[d + 1] = Oacc[dn8][1];
|
||||
}
|
||||
if (r1 < G) {
|
||||
int h = q_head0 + r1;
|
||||
float* op = p.o_part + split_slot(h) * HEAD_DIM;
|
||||
float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
|
||||
op[d] = Oacc[dn8][2];
|
||||
op[d + 1] = Oacc[dn8][3];
|
||||
}
|
||||
|
||||
@@ -0,0 +1,195 @@
|
||||
#pragma once
|
||||
// Shared attention dispatchers — used by both production .cu and test .cu.
|
||||
// No torch dependency; pure CUDA.
|
||||
|
||||
#include <cuda_runtime.h>
|
||||
#include <algorithm>
|
||||
#include "attn_warp_utils.cuh"
|
||||
#include "attn_prefill_split_q.cuh"
|
||||
#include "attn_decode_split_kv.cuh"
|
||||
#include "attn_paged_decode_split_kv.cuh"
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
#include "attn_prefill_split_q_mma.cuh"
|
||||
#include "attn_decode_split_kv_mma.cuh"
|
||||
#include "attn_paged_decode_split_kv_mma.cuh"
|
||||
#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) {
|
||||
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)));
|
||||
}
|
||||
|
||||
// ======================================================================
|
||||
// Prefill
|
||||
// ======================================================================
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||
static inline void launch_prefill_mma(AttentionParams<bf16>& p) {
|
||||
constexpr int WARPS = 4;
|
||||
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
|
||||
using Traits = KernelTraits<HEAD_DIM, BC, WARPS, 2>;
|
||||
dim3 grid((p.q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS), p.q_head, p.batch);
|
||||
dim3 block(Traits::NUM_THREADS);
|
||||
attn_prefill_split_q_mma_kernel<Traits, IsCausal, HasMask><<<grid, block>>>(p);
|
||||
}
|
||||
#endif
|
||||
|
||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||
static inline void launch_prefill_scalar(AttentionParams<bf16>& p) {
|
||||
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
||||
dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
|
||||
dim3 block(G, ROWS);
|
||||
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask><<<grid, block>>>(p);
|
||||
}
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static inline void dispatch_prefill(AttentionParams<bf16>& p) {
|
||||
bool is_causal = (p.causal_offset >= 0);
|
||||
bool has_mask = (p.use_mask && p.mask);
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
if (is_causal) {
|
||||
if (has_mask) launch_prefill_mma<HEAD_DIM, true, true>(p);
|
||||
else launch_prefill_mma<HEAD_DIM, true, false>(p);
|
||||
} else {
|
||||
if (has_mask) launch_prefill_mma<HEAD_DIM, false, true>(p);
|
||||
else launch_prefill_mma<HEAD_DIM, false, false>(p);
|
||||
}
|
||||
#else
|
||||
if (is_causal) {
|
||||
if (has_mask) launch_prefill_scalar<HEAD_DIM, true, true>(p);
|
||||
else launch_prefill_scalar<HEAD_DIM, true, false>(p);
|
||||
} else {
|
||||
if (has_mask) launch_prefill_scalar<HEAD_DIM, false, true>(p);
|
||||
else launch_prefill_scalar<HEAD_DIM, false, false>(p);
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
// ======================================================================
|
||||
// Decode
|
||||
// ======================================================================
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
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>;
|
||||
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
|
||||
attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask><<<grid, 32>>>(p);
|
||||
}
|
||||
#endif
|
||||
|
||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||
static inline void launch_decode_scalar(AttentionParams<bf16>& p, int group_size) {
|
||||
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
||||
int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit
|
||||
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
|
||||
dim3 block(32, g);
|
||||
attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
|
||||
}
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static inline void dispatch_decode(AttentionParams<bf16>& p) {
|
||||
bool is_causal = (p.causal_offset >= 0);
|
||||
bool has_mask = (p.use_mask && p.mask);
|
||||
int group_size = p.q_head / p.kv_head;
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
if (is_causal) {
|
||||
if (has_mask) launch_decode_mma<HEAD_DIM, true, true>(p, group_size);
|
||||
else launch_decode_mma<HEAD_DIM, true, false>(p, group_size);
|
||||
} else {
|
||||
if (has_mask) launch_decode_mma<HEAD_DIM, false, true>(p, group_size);
|
||||
else launch_decode_mma<HEAD_DIM, false, false>(p, group_size);
|
||||
}
|
||||
#else
|
||||
if (is_causal) {
|
||||
if (has_mask) launch_decode_scalar<HEAD_DIM, true, true>(p, group_size);
|
||||
else launch_decode_scalar<HEAD_DIM, true, false>(p, group_size);
|
||||
} else {
|
||||
if (has_mask) launch_decode_scalar<HEAD_DIM, false, true>(p, group_size);
|
||||
else launch_decode_scalar<HEAD_DIM, false, false>(p, group_size);
|
||||
}
|
||||
#endif
|
||||
|
||||
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
|
||||
// ======================================================================
|
||||
// Paged Decode
|
||||
// ======================================================================
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
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);
|
||||
}
|
||||
}
|
||||
#endif
|
||||
|
||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||
static inline void launch_paged_decode_scalar(PagedAttentionParams<bf16>& p, int group_size) {
|
||||
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);
|
||||
int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit
|
||||
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
|
||||
dim3 block(32, g);
|
||||
paged_attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
|
||||
}
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static inline void dispatch_paged_decode(PagedAttentionParams<bf16>& p) {
|
||||
bool is_causal = (p.causal_offset >= 0);
|
||||
bool has_mask = (p.use_mask && p.mask);
|
||||
int group_size = p.q_head / p.kv_head;
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
if (is_causal) {
|
||||
if (has_mask) launch_paged_decode_mma<HEAD_DIM, true, true>(p, group_size);
|
||||
else launch_paged_decode_mma<HEAD_DIM, true, false>(p, group_size);
|
||||
} else {
|
||||
if (has_mask) launch_paged_decode_mma<HEAD_DIM, false, true>(p, group_size);
|
||||
else launch_paged_decode_mma<HEAD_DIM, false, false>(p, group_size);
|
||||
}
|
||||
#else
|
||||
if (is_causal) {
|
||||
if (has_mask) launch_paged_decode_scalar<HEAD_DIM, true, true>(p, group_size);
|
||||
else launch_paged_decode_scalar<HEAD_DIM, true, false>(p, group_size);
|
||||
} else {
|
||||
if (has_mask) launch_paged_decode_scalar<HEAD_DIM, false, true>(p, group_size);
|
||||
else launch_paged_decode_scalar<HEAD_DIM, false, false>(p, group_size);
|
||||
}
|
||||
#endif
|
||||
|
||||
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
@@ -2,16 +2,10 @@
|
||||
#include <torch/extension.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include "attn_common.h"
|
||||
#include "attn_warp_utils.cuh"
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
|
||||
inline int compute_num_splits(int base_blocks, int tiles_total) {
|
||||
int sm_count = 0;
|
||||
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
|
||||
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
|
||||
return std::max(1, std::min(n, std::min(tiles_total, 32)));
|
||||
}
|
||||
|
||||
// Dispatch head_dim: shared macro — avoids C++20 lambda template syntax.
|
||||
// Usage: DISPATCH_HEAD_DIM(hd, fn, arg)
|
||||
// Expands to: fn<32>(arg); fn<64>(arg); etc.
|
||||
@@ -29,8 +23,8 @@ inline int compute_num_splits(int base_blocks, int tiles_total) {
|
||||
template<typename P>
|
||||
inline void alloc_split_partials(P& p) {
|
||||
auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
|
||||
auto o_part = torch::empty({p.batch, p.q_head, p.num_splits, p.head_dim}, fopt);
|
||||
auto ml_part = torch::empty({p.batch, p.q_head, p.num_splits, 2}, fopt);
|
||||
auto o_part = torch::empty({p.batch, p.q_head, MAX_SPLITS, p.head_dim}, fopt);
|
||||
auto ml_part = torch::empty({p.batch, p.q_head, MAX_SPLITS, 2}, fopt);
|
||||
p.o_part = (float*)o_part.data_ptr();
|
||||
p.ml_part = (float*)ml_part.data_ptr();
|
||||
}
|
||||
|
||||
@@ -3,10 +3,41 @@
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
// Shared MMA utilities for tensor-core GQA kernels.
|
||||
// mma.sync.m16n8k16 PTX wrappers, ldmatrix helpers, and bf16 packing.
|
||||
// ============================================================================
|
||||
// KernelTraits — FlashAttention-v2 style compile-time configuration bundle.
|
||||
//
|
||||
// Bundles all dimension-dependent constants so device functions only need a
|
||||
// single Traits template parameter rather than scattered <KD, NC8, KT2, ...>.
|
||||
// ============================================================================
|
||||
template <int HEAD_DIM_, int BC_, int WARPS_, int STAGES_>
|
||||
struct KernelTraits {
|
||||
static constexpr int HEAD_DIM = HEAD_DIM_;
|
||||
static constexpr int BC = BC_; // K/V tile size along seq dim
|
||||
static constexpr int WARPS = WARPS_; // warps per block
|
||||
static constexpr int STAGES = STAGES_; // double-buffer stages (1 or 2)
|
||||
|
||||
static constexpr int BR = 16; // Q rows per warp (mma M=16)
|
||||
|
||||
// Derived: mma.sync.m16n8k16 tile counts
|
||||
static constexpr int KD = HEAD_DIM / 16; // Q/K k-slides
|
||||
static constexpr int NC8 = BC / 8; // S n-tiles (N=8)
|
||||
static constexpr int KT2 = BC / 16; // P k-tiles (K=16)
|
||||
static constexpr int DN8 = HEAD_DIM / 8; // O n-tiles (N=8)
|
||||
|
||||
static constexpr int LD = HEAD_DIM; // smem leading dim
|
||||
|
||||
// XOR swizzle chunk bits for ldmatrix bank-conflict avoidance.
|
||||
// mask = log2(LD/8) bits, clamped to stay within LD.
|
||||
static constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
|
||||
|
||||
static constexpr int NUM_THREADS = WARPS * 32;
|
||||
static constexpr int VEC = 8; // bf16 per cp.async unit (16 bytes)
|
||||
static constexpr int TOTAL = BC * HEAD_DIM; // total elements per tile
|
||||
};
|
||||
|
||||
// ---- PTX wrappers ----
|
||||
using bf16 = __nv_bfloat16;
|
||||
|
||||
// mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32
|
||||
__device__ __forceinline__ void mma16816(float* d, const unsigned* a,
|
||||
const unsigned* b, const float* c) {
|
||||
asm volatile(
|
||||
@@ -37,9 +68,7 @@ __device__ __forceinline__ unsigned pkb(bf16 a, bf16 b) {
|
||||
}
|
||||
|
||||
// ldmatrix: cooperatively load mma fragments from smem (one instruction per
|
||||
// 16x16 / 16x8 tile) with the exact register layout mma expects — replaces the
|
||||
// scalar per-thread fragment packing, cutting shared-load instructions and bank
|
||||
// conflicts. Each lane supplies the shared address of one 8-wide row.
|
||||
// 16x16 / 16x8 tile) with the exact register layout mma expects.
|
||||
__device__ __forceinline__ void ldmatrix_x4(unsigned* r, const bf16* p) {
|
||||
unsigned a = __cvta_generic_to_shared(p);
|
||||
asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];"
|
||||
@@ -60,29 +89,19 @@ __device__ __forceinline__ void ldmatrix_x2_trans(unsigned* r, const bf16* p) {
|
||||
}
|
||||
|
||||
// XOR swizzle for shared-memory column at 8-bf16 chunk granularity.
|
||||
// Eliminates ldmatrix bank conflicts without LD padding: consecutive rows
|
||||
// land in distinct bank groups. swiz_col(d, r, mask) = ((d>>3)^(r&mask))<<3 | (d&7).
|
||||
// mask must cover log2(HEAD_DIM/8) chunk bits but stay within LD: use 7 for
|
||||
// HEAD_DIM>=64 (8+ chunks), 3 for HEAD_DIM=32 (4 chunks). Default 7 keeps
|
||||
// existing HEAD_DIM>=64 call sites working unchanged.
|
||||
__device__ __forceinline__ int swiz_col(int d, int r, int mask = 7) {
|
||||
return ((d >> 3) ^ (r & mask)) << 3 | (d & 7);
|
||||
}
|
||||
|
||||
// cp.async: copy 16 bytes (8 bf16) from global to shared memory directly,
|
||||
// bypassing registers. Eliminates shared-store bank conflicts and cuts
|
||||
// load-loop instruction count in half (1 cp.async vs 1 LDG + 1 STS).
|
||||
// Requires sm_80+.
|
||||
// cp.async: copy 16 bytes (8 bf16) from global to shared memory directly.
|
||||
__device__ __forceinline__ void cp_async_16(bf16* smem_ptr, const void* gmem_ptr) {
|
||||
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
|
||||
asm volatile("cp.async.ca.shared.global [%0], [%1], 16;"
|
||||
:: "r"(smem_addr), "l"(gmem_ptr));
|
||||
}
|
||||
|
||||
// Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill the
|
||||
// destination (src-size operand = 0 → no bytes read from src, so an
|
||||
// out-of-bounds src address is never dereferenced). Lets full and partial
|
||||
// tiles share one uniform async load path — no scalar fallback branch.
|
||||
// Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill.
|
||||
// src_size=0 → no bytes read from src, so out-of-bounds src address is safe.
|
||||
__device__ __forceinline__ void cp_async_16_pred(bf16* smem_ptr,
|
||||
const void* gmem_ptr,
|
||||
bool pred) {
|
||||
@@ -100,9 +119,6 @@ __device__ __forceinline__ void cp_async_wait_all() {
|
||||
asm volatile("cp.async.wait_all;");
|
||||
}
|
||||
|
||||
// Wait until at most N commit groups are still in flight. Used for
|
||||
// double-buffered pipelining: wait_group<1> lets the next tile's cp.async
|
||||
// continue while ensuring the current tile's data is ready.
|
||||
template <int N>
|
||||
__device__ __forceinline__ void cp_async_wait_group() {
|
||||
asm volatile("cp.async.wait_group %0;" :: "n"(N));
|
||||
@@ -139,78 +155,65 @@ __device__ inline void load_q_mma_frags(
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Shared MMA compute functions — used by both decode and prefill MMA kernels.
|
||||
// Extracted because S=Q@K^T, online softmax, and P@V are structurally identical
|
||||
// between the two kernels; only the per-row causal/mask bounds differ.
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// S = Q @ K^T (Qa pre-loaded by the caller; scale applied post-mma in the
|
||||
// caller to avoid bf16 precision loss).
|
||||
// LD and SWIZ_MASK are constexpr in the calling kernel — passing them as
|
||||
// runtime ints lets the compiler fold them while keeping the signature clean.
|
||||
template <int KD, int NC8>
|
||||
// Traits provides KD, NC8, LD, and SWIZ_MASK.
|
||||
// ---------------------------------------------------------------------------
|
||||
template <typename Traits>
|
||||
__device__ inline void mma_compute_scores(
|
||||
const unsigned Qa[KD][4],
|
||||
const unsigned Qa[Traits::KD][4],
|
||||
const bf16* __restrict__ sK,
|
||||
int LD,
|
||||
int SWIZ_MASK,
|
||||
int lane,
|
||||
float Sacc[NC8][4])
|
||||
float Sacc[Traits::NC8][4])
|
||||
{
|
||||
#pragma unroll
|
||||
for (int n8 = 0; n8 < NC8; n8++) {
|
||||
for (int n8 = 0; n8 < Traits::NC8; n8++) {
|
||||
Sacc[n8][0] = Sacc[n8][1] = Sacc[n8][2] = Sacc[n8][3] = 0.0f;
|
||||
int krow_l = n8 * 8 + (lane & 7);
|
||||
int kcol_h = (lane & 8) ? 8 : 0;
|
||||
#pragma unroll
|
||||
for (int kt = 0; kt < KD; kt++) {
|
||||
for (int kt = 0; kt < Traits::KD; kt++) {
|
||||
unsigned b[2];
|
||||
ldmatrix_x2(b, &sK[krow_l * LD + swiz_col(kt * 16 + kcol_h, krow_l, SWIZ_MASK)]);
|
||||
ldmatrix_x2(b, &sK[krow_l * Traits::LD
|
||||
+ swiz_col(kt * 16 + kcol_h, krow_l, Traits::SWIZ_MASK)]);
|
||||
mma16816(Sacc[n8], Qa[kt], b, Sacc[n8]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Online softmax + Oacc rescale for one K/V tile.
|
||||
// maxc0/maxc1: per-row KV column bounds (prefill: per-query-row causal limits;
|
||||
// decode: same value for both rows since q_len==1).
|
||||
// qrow0/qrow1: query row indices (for 3D mask indexing; decode passes 0).
|
||||
// mask_b_stride/mask_q_stride: mask layout (2D: mask_q_stride=0; 3D: =kv_len).
|
||||
// Reads Sacc (Q@K^T scores), applies causal/mask, computes P = exp(S - nm),
|
||||
// rescales Oacc by exp(m_old - nm), and updates m/l — all in place.
|
||||
template <int NC8, int DN8>
|
||||
//
|
||||
// HasMask is a compile-time template bool: when false, the mask branch is
|
||||
// entirely dead-code-eliminated from the inner unrolled loop.
|
||||
// ---------------------------------------------------------------------------
|
||||
template <typename Traits, bool HasMask>
|
||||
__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 maxc0, int maxc1,
|
||||
int qrow0, int qrow1,
|
||||
int mask_b_stride, int mask_q_stride,
|
||||
int mask_batch,
|
||||
const bool* __restrict__ mask,
|
||||
bool has_mask,
|
||||
float Sacc[NC8][4],
|
||||
float Oacc[DN8][4],
|
||||
float Sacc[Traits::NC8][4],
|
||||
float Oacc[Traits::DN8][4],
|
||||
float& m0, float& m1,
|
||||
float& l0, float& l1,
|
||||
int lane)
|
||||
{
|
||||
int tid4 = lane & 3;
|
||||
|
||||
// Mask out-of-bounds / masked columns: set -FLT_MAX so expf → 0 downstream
|
||||
// without per-element sentinel checks. Compute tile-local row maxima.
|
||||
float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX;
|
||||
int mask_base0 = mask_batch * mask_b_stride + qrow0 * mask_q_stride;
|
||||
int mask_base1 = mask_batch * mask_b_stride + qrow1 * mask_q_stride;
|
||||
#pragma unroll
|
||||
for (int n8 = 0; n8 < NC8; n8++) {
|
||||
for (int n8 = 0; n8 < Traits::NC8; n8++) {
|
||||
int cc = kv0 + n8 * 8 + 2 * tid4;
|
||||
int c1 = cc + 1;
|
||||
bool b0 = (cc >= maxc0) || (has_mask && !mask[mask_base0 + cc]);
|
||||
bool b1 = (c1 >= maxc0) || (has_mask && !mask[mask_base0 + c1]);
|
||||
bool b2 = (cc >= maxc1) || (has_mask && !mask[mask_base1 + cc]);
|
||||
bool b3 = (c1 >= maxc1) || (has_mask && !mask[mask_base1 + c1]);
|
||||
bool b0 = (cc >= maxc0) || (HasMask && !mask[mask_base0 + cc]);
|
||||
bool b1 = (c1 >= maxc0) || (HasMask && !mask[mask_base0 + c1]);
|
||||
bool b2 = (cc >= maxc1) || (HasMask && !mask[mask_base1 + cc]);
|
||||
bool b3 = (c1 >= maxc1) || (HasMask && !mask[mask_base1 + c1]);
|
||||
float s0 = b0 ? -FLT_MAX : Sacc[n8][0];
|
||||
float s1 = b1 ? -FLT_MAX : Sacc[n8][1];
|
||||
float s2 = b2 ? -FLT_MAX : Sacc[n8][2];
|
||||
@@ -220,29 +223,20 @@ __device__ inline void mma_softmax_tile(
|
||||
rmax0 = fmaxf(rmax0, fmaxf(s0, s1));
|
||||
rmax1 = fmaxf(rmax1, fmaxf(s2, s3));
|
||||
}
|
||||
// Warp-reduce row maxima across the 4-lane thread group (xor 1, xor 2).
|
||||
rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 1));
|
||||
rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 2));
|
||||
rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 1));
|
||||
rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 2));
|
||||
|
||||
// nm = max(running max m, tile-local max rmax) — updated running maximum.
|
||||
float nm0 = fmaxf(m0, rmax0), nm1 = fmaxf(m1, rmax1);
|
||||
// corr rescales Oacc and l by exp(m_old - nm). When all-masked (m == nm ==
|
||||
// -FLT_MAX), exp(0) = 1 — correct, no guard needed.
|
||||
float corr0 = __expf(m0 - nm0);
|
||||
float corr1 = __expf(m1 - nm1);
|
||||
// pn guards only the all-masked-row edge: if nm == -FLT_MAX, exp(S - nm)
|
||||
// gives 1 not 0 for masked entries. Two scalar masks replace 4*NC8
|
||||
// per-element comparisons.
|
||||
float pn0 = (nm0 == -FLT_MAX) ? 0.0f : 1.0f;
|
||||
float pn1 = (nm1 == -FLT_MAX) ? 0.0f : 1.0f;
|
||||
|
||||
// P = exp(S - nm) for each element. Masked entries (Sacc = -FLT_MAX) give
|
||||
// exp(-inf) ≈ 0 naturally; pn zero-fills the all-masked-row edge.
|
||||
float rsum0 = 0.0f, rsum1 = 0.0f;
|
||||
#pragma unroll
|
||||
for (int n8 = 0; n8 < NC8; n8++) {
|
||||
for (int n8 = 0; n8 < Traits::NC8; n8++) {
|
||||
float p0 = pn0 * __expf(Sacc[n8][0] - nm0);
|
||||
float p1 = pn0 * __expf(Sacc[n8][1] - nm0);
|
||||
float p2 = pn1 * __expf(Sacc[n8][2] - nm1);
|
||||
@@ -261,22 +255,25 @@ __device__ inline void mma_softmax_tile(
|
||||
m0 = nm0; m1 = nm1;
|
||||
|
||||
#pragma unroll
|
||||
for (int j = 0; j < DN8; j++) {
|
||||
for (int j = 0; j < Traits::DN8; j++) {
|
||||
Oacc[j][0] *= corr0; Oacc[j][1] *= corr0;
|
||||
Oacc[j][2] *= corr1; Oacc[j][3] *= corr1;
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// O += P @ V (Sacc must contain P = attention weights after softmax).
|
||||
template <int DN8, int KT2>
|
||||
// Traits provides DN8, KT2, LD, and SWIZ_MASK.
|
||||
// ---------------------------------------------------------------------------
|
||||
template <typename Traits>
|
||||
__device__ inline void mma_pv_accumulate(
|
||||
float Sacc[][4],
|
||||
const bf16* __restrict__ sV,
|
||||
int LD, int SWIZ_MASK, int lane,
|
||||
float Oacc[DN8][4])
|
||||
int lane,
|
||||
float Oacc[Traits::DN8][4])
|
||||
{
|
||||
#pragma unroll
|
||||
for (int kt2 = 0; kt2 < KT2; kt2++) {
|
||||
for (int kt2 = 0; kt2 < Traits::KT2; kt2++) {
|
||||
unsigned Pa[4];
|
||||
Pa[0] = pk2(Sacc[kt2 * 2][0], Sacc[kt2 * 2][1]);
|
||||
Pa[1] = pk2(Sacc[kt2 * 2][2], Sacc[kt2 * 2][3]);
|
||||
@@ -284,9 +281,10 @@ __device__ inline void mma_pv_accumulate(
|
||||
Pa[3] = pk2(Sacc[kt2 * 2 + 1][2], Sacc[kt2 * 2 + 1][3]);
|
||||
int vrow_l = kt2 * 16 + (lane & 15);
|
||||
#pragma unroll
|
||||
for (int dn8 = 0; dn8 < DN8; dn8++) {
|
||||
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
|
||||
unsigned b[2];
|
||||
ldmatrix_x2_trans(b, &sV[vrow_l * LD + swiz_col(dn8 * 8, vrow_l, SWIZ_MASK)]);
|
||||
ldmatrix_x2_trans(b, &sV[vrow_l * Traits::LD
|
||||
+ swiz_col(dn8 * 8, vrow_l, Traits::SWIZ_MASK)]);
|
||||
mma16816(Oacc[dn8], Pa, b, Oacc[dn8]);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,47 +1,6 @@
|
||||
#include "attn_paged_decode_split_kv.cuh"
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
#include "attn_paged_decode_split_kv_mma.cuh"
|
||||
#endif
|
||||
|
||||
#include "attn_dispatchers.cuh"
|
||||
#include "attn_entry_utils.cuh"
|
||||
|
||||
static void launch_paged_scalar_decode(PagedAttentionParams<bf16>& p) {
|
||||
int group_size = p.q_head / p.kv_head;
|
||||
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||
alloc_split_partials(p);
|
||||
|
||||
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
|
||||
dim3 grid = dim3(p.batch * p.kv_head, 1, p.num_splits);
|
||||
dim3 block = dim3(32, group_size);
|
||||
paged_attn_decode_split_kv_kernel<<<grid, block, smem>>>(p);
|
||||
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
|
||||
static void launch_paged_mma_decode(PagedAttentionParams<bf16>& p) {
|
||||
int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
||||
alloc_split_partials(p);
|
||||
|
||||
paged_attn_decode_split_kv_mma_kernel<HEAD_DIM, BC, STAGES><<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
|
||||
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
#endif
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static void dispatch_paged_decode(PagedAttentionParams<bf16>& p) {
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
int G = p.q_head / p.kv_head;
|
||||
if (G >= 1 && G <= 16 && p.page_size >= 32) {
|
||||
launch_paged_mma_decode<HEAD_DIM, 32>(p);
|
||||
return;
|
||||
}
|
||||
#endif
|
||||
launch_paged_scalar_decode(p);
|
||||
}
|
||||
|
||||
torch::Tensor attn_paged_decode(
|
||||
torch::Tensor q,
|
||||
torch::Tensor page_table,
|
||||
@@ -62,6 +21,7 @@ torch::Tensor attn_paged_decode(
|
||||
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
|
||||
p.o = (bf16*)O_view.data_ptr();
|
||||
|
||||
alloc_split_partials(p);
|
||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p);
|
||||
return O;
|
||||
}
|
||||
|
||||
@@ -2,17 +2,10 @@
|
||||
#include <cuda_bf16.h>
|
||||
#include <float.h>
|
||||
#include "attn_common.h"
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
#include "attn_warp_utils.cuh"
|
||||
constexpr int PDC_CHUNK = 64;
|
||||
|
||||
__device__ inline float paged_warp_reduce_sum(float val) {
|
||||
for (int offset = 16; offset > 0; offset >>= 1)
|
||||
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
|
||||
return val;
|
||||
}
|
||||
|
||||
// Split-KV scalar decode: one warp per query head, grid.z partitions KV.
|
||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||
__global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p) {
|
||||
int batch = blockIdx.x / p.kv_head;
|
||||
int kv_head = blockIdx.x % p.kv_head;
|
||||
@@ -22,7 +15,6 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
|
||||
int lane = threadIdx.x;
|
||||
int hd_per_thread = p.head_dim / 32;
|
||||
|
||||
// Q: stride-based [batch, q_head, q_len=1, head_dim]
|
||||
float q_reg[8];
|
||||
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
||||
+ lane * hd_per_thread * p.q_stride_d;
|
||||
@@ -46,7 +38,8 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
|
||||
int this_chunk = min(PDC_CHUNK, p.kv_len - chunk_start);
|
||||
|
||||
int total = this_chunk * p.head_dim;
|
||||
for (int i = threadIdx.y * 32 + lane; i < total; i += blockDim.x * blockDim.y) {
|
||||
for (int i = threadIdx.y * 32 + lane; i < total;
|
||||
i += blockDim.x * blockDim.y) {
|
||||
int s = i / p.head_dim;
|
||||
int d_dim = i % p.head_dim;
|
||||
int pos = chunk_start + s;
|
||||
@@ -69,14 +62,19 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
|
||||
float partial = 0.0f;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
partial += q_reg[i] * __bfloat162float(k_smem[s * p.head_dim + lane * hd_per_thread + i]);
|
||||
partial = paged_warp_reduce_sum(partial) * p.scale;
|
||||
partial += q_reg[i] * __bfloat162float(
|
||||
k_smem[s * p.head_dim + lane * hd_per_thread + i]);
|
||||
partial = warp_reduce_sum(partial) * p.scale;
|
||||
|
||||
int kv_idx = chunk_start + s;
|
||||
if (p.use_mask && p.mask && !p.mask[mask_base + kv_idx])
|
||||
partial = -FLT_MAX;
|
||||
if (p.causal_offset >= 0 && kv_idx > p.causal_offset)
|
||||
partial = -FLT_MAX;
|
||||
if constexpr (HasMask) {
|
||||
if (!p.mask[mask_base + kv_idx])
|
||||
partial = -FLT_MAX;
|
||||
}
|
||||
if constexpr (IsCausal) {
|
||||
if (kv_idx > p.causal_offset)
|
||||
partial = -FLT_MAX;
|
||||
}
|
||||
|
||||
float new_m = fmaxf(m, partial);
|
||||
float alpha = expf(m - new_m);
|
||||
@@ -93,11 +91,12 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
|
||||
+ (int64_t)kv_head * p.head_dim;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
acc_reg[i] = acc_reg[i] * alpha + __bfloat162float(p.v_cache[v_base + lane * hd_per_thread + i]) * beta;
|
||||
acc_reg[i] = fmaf(acc_reg[i], alpha,
|
||||
__bfloat162float(p.v_cache[v_base + lane * hd_per_thread + i]) * beta);
|
||||
} else {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
acc_reg[i] = acc_reg[i] * alpha + 0.0f * beta;
|
||||
acc_reg[i] = fmaf(acc_reg[i], alpha, 0.0f);
|
||||
}
|
||||
m = new_m;
|
||||
}
|
||||
@@ -105,7 +104,7 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
|
||||
}
|
||||
|
||||
size_t bh = (size_t)batch * p.q_head + q_head;
|
||||
size_t slot = bh * p.num_splits + split;
|
||||
size_t slot = bh * MAX_SPLITS + split;
|
||||
int d0 = lane * hd_per_thread;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
@@ -124,7 +123,7 @@ __global__ void paged_attn_decode_combine_kernel(PagedAttentionParams<bf16> p) {
|
||||
int batch = bh / p.q_head;
|
||||
int q_head = bh % p.q_head;
|
||||
|
||||
size_t split_base = (size_t)bh * p.num_splits;
|
||||
size_t split_base = (size_t)bh * MAX_SPLITS;
|
||||
const float* mlp = p.ml_part + split_base * 2;
|
||||
const float* op = p.o_part + split_base * p.head_dim;
|
||||
|
||||
@@ -134,10 +133,10 @@ __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);
|
||||
acc = acc * corr + op[s * p.head_dim + d] * e;
|
||||
l = l * corr + li * e;
|
||||
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;
|
||||
}
|
||||
|
||||
|
||||
@@ -3,153 +3,144 @@
|
||||
#include <cuda_bf16.h>
|
||||
#include "attn_common.h"
|
||||
#include "attn_mma_utils.cuh"
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
#include "attn_warp_utils.cuh"
|
||||
|
||||
// Paged split-KV tensor-core decode via GQA head-packing.
|
||||
// Identical algorithm to attn_decode_split_kv_mma_kernel but reads K/V
|
||||
// directly from the page pool through a page table, eliminating the gather
|
||||
// copy. Each tile (BC=32) fits within a single page (page_size >= 32), so
|
||||
// the page-table lookup happens once per tile for cp.async.
|
||||
|
||||
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
|
||||
// Reads K/V directly from the page pool through a page table — one tile
|
||||
// (BC=32) fits within a single page (page_size >= 32), so the page-table
|
||||
// lookup happens once per tile for cp.async.
|
||||
//
|
||||
// IsCausal and HasMask are compile-time bools.
|
||||
template <typename Traits, bool IsCausal, bool HasMask>
|
||||
__global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16> p) {
|
||||
constexpr int KD = HEAD_DIM / 16;
|
||||
constexpr int NC8 = BC / 8;
|
||||
constexpr int KT2 = BC / 16;
|
||||
constexpr int DN8 = HEAD_DIM / 8;
|
||||
constexpr int LD = HEAD_DIM;
|
||||
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
|
||||
constexpr int VEC = 8;
|
||||
constexpr int TOTAL = BC * HEAD_DIM;
|
||||
|
||||
const int lane = threadIdx.x;
|
||||
const int gid = lane >> 2;
|
||||
const int tid4 = lane & 3;
|
||||
|
||||
const int kv_head = blockIdx.x;
|
||||
const int pass = blockIdx.x / p.kv_head;
|
||||
const int kv_head = blockIdx.x % p.kv_head;
|
||||
const int batch = blockIdx.y;
|
||||
const int split = blockIdx.z;
|
||||
const int G = p.q_head / p.kv_head;
|
||||
const int q_head0 = kv_head * G;
|
||||
|
||||
__shared__ __align__(16) bf16 sK[STAGES * BC * LD];
|
||||
__shared__ __align__(16) bf16 sV[STAGES * BC * LD];
|
||||
constexpr int MAX_G = 16;
|
||||
const int G_total = p.q_head / p.kv_head;
|
||||
const int g_begin = pass * MAX_G;
|
||||
const int G = min(MAX_G, G_total - g_begin);
|
||||
const int q_head0 = kv_head * G_total + g_begin;
|
||||
|
||||
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
|
||||
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
|
||||
|
||||
// ---- Load Q directly from global into mma A-operand registers ----
|
||||
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
|
||||
const int qra = gid;
|
||||
const int qrb = gid + 8;
|
||||
const bool va = qra < G, vb = qrb < G;
|
||||
unsigned Qa[KD][4];
|
||||
load_q_mma_frags<KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
|
||||
qra, qrb, va, vb, tid4, Qa);
|
||||
unsigned Qa[Traits::KD][4];
|
||||
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
|
||||
qra, qrb, va, vb, tid4, Qa);
|
||||
|
||||
float Oacc[DN8][4];
|
||||
#pragma unroll
|
||||
for (int j = 0; j < DN8; j++)
|
||||
float Oacc[Traits::DN8][4];
|
||||
#pragma unroll
|
||||
for (int j = 0; j < Traits::DN8; j++)
|
||||
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
||||
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
||||
|
||||
const int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||
const int tiles_total = (p.kv_len + Traits::BC - 1) / Traits::BC;
|
||||
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
|
||||
const int ti_begin = split * tiles_per_split;
|
||||
const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
|
||||
const int has_mask = p.use_mask && p.mask;
|
||||
|
||||
// Paged strides (constant for the block)
|
||||
const int64_t page_stride = (int64_t)p.page_size * p.kv_head * HEAD_DIM;
|
||||
const int64_t pos_stride = (int64_t)p.kv_head * HEAD_DIM;
|
||||
const int64_t head_off = (int64_t)kv_head * HEAD_DIM;
|
||||
const int64_t page_stride = (int64_t)p.page_size * p.kv_head * Traits::HEAD_DIM;
|
||||
const int64_t pos_stride = (int64_t)p.kv_head * Traits::HEAD_DIM;
|
||||
const int64_t head_off = (int64_t)kv_head * Traits::HEAD_DIM;
|
||||
|
||||
// ---- Load tile lambda: predicated cp.async, paged addressing ----
|
||||
// ---- Load tile lambda: paged addressing ----
|
||||
auto load_tile = [&](int ti, int buf) {
|
||||
int kv0 = ti * BC;
|
||||
bf16* dK = sK + buf * BC * LD;
|
||||
bf16* dV = sV + buf * BC * LD;
|
||||
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 * VEC; i < TOTAL; i += 32 * VEC) {
|
||||
int r = i / HEAD_DIM, d = i % HEAD_DIM;
|
||||
#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;
|
||||
int page_off = kc % p.page_size;
|
||||
int64_t gmem_base = (int64_t)phys_page * page_stride
|
||||
+ (int64_t)page_off * pos_stride
|
||||
+ head_off;
|
||||
int off = r * LD + swiz_col(d, r, SWIZ_MASK);
|
||||
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
|
||||
cp_async_16_pred(&dK[off], &p.k_cache[gmem_base + d], valid);
|
||||
cp_async_16_pred(&dV[off], &p.v_cache[gmem_base + d], valid);
|
||||
}
|
||||
cp_async_commit();
|
||||
};
|
||||
|
||||
// ---- Prologue: issue first tile load ----
|
||||
constexpr int BUF_MASK = (Traits::STAGES > 1) ? (Traits::STAGES - 1) : 0;
|
||||
|
||||
if (ti_begin < ti_end) {
|
||||
load_tile(ti_begin, 0);
|
||||
}
|
||||
|
||||
for (int ti = ti_begin; ti < ti_end; ti++) {
|
||||
constexpr int BUF_MASK = (STAGES > 1) ? (STAGES - 1) : 0;
|
||||
int buf = (ti - ti_begin) & BUF_MASK;
|
||||
|
||||
cp_async_wait_group<0>();
|
||||
__syncwarp();
|
||||
if constexpr (STAGES > 1) {
|
||||
if constexpr (Traits::STAGES > 1) {
|
||||
if (ti + 1 < ti_end)
|
||||
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
|
||||
}
|
||||
|
||||
const bf16* bK = sK + buf * BC * LD;
|
||||
const bf16* bV = sV + buf * BC * LD;
|
||||
int kv0 = ti * BC;
|
||||
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
|
||||
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
|
||||
int kv0 = ti * Traits::BC;
|
||||
|
||||
float Sacc[NC8][4];
|
||||
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
|
||||
float Sacc[Traits::NC8][4];
|
||||
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
|
||||
|
||||
#pragma unroll
|
||||
for (int n8 = 0; n8 < NC8; n8++)
|
||||
#pragma unroll
|
||||
for (int n8 = 0; n8 < Traits::NC8; n8++)
|
||||
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||
|
||||
// Decode: q_len=1, so qrow0=qrow1=0, mask_q_stride irrelevant
|
||||
int maxc = (p.causal_offset >= 0) ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
||||
mma_softmax_tile<NC8, DN8>(kv0, maxc, maxc,
|
||||
0, 0,
|
||||
p.mask_b_stride, 0,
|
||||
batch,
|
||||
p.mask, has_mask,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
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,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
|
||||
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc);
|
||||
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
||||
__syncwarp();
|
||||
|
||||
if constexpr (STAGES == 1) {
|
||||
if constexpr (Traits::STAGES == 1) {
|
||||
if (ti + 1 < ti_end)
|
||||
load_tile(ti + 1, 0);
|
||||
}
|
||||
}
|
||||
|
||||
// ---- write UN-normalised partials for this split ----
|
||||
auto split_slot = [&](int h) -> size_t {
|
||||
size_t bh = (size_t)batch * p.q_head + h;
|
||||
return bh * p.num_splits + split;
|
||||
return bh * MAX_SPLITS + split;
|
||||
};
|
||||
#pragma unroll
|
||||
for (int dn8 = 0; dn8 < DN8; dn8++) {
|
||||
#pragma unroll
|
||||
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
|
||||
int d = dn8 * 8 + 2 * tid4;
|
||||
int r0 = gid, r1 = gid + 8;
|
||||
if (r0 < G) {
|
||||
int h = q_head0 + r0;
|
||||
float* op = p.o_part + split_slot(h) * HEAD_DIM;
|
||||
float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
|
||||
op[d] = Oacc[dn8][0];
|
||||
op[d + 1] = Oacc[dn8][1];
|
||||
}
|
||||
if (r1 < G) {
|
||||
int h = q_head0 + r1;
|
||||
float* op = p.o_part + split_slot(h) * HEAD_DIM;
|
||||
float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
|
||||
op[d] = Oacc[dn8][2];
|
||||
op[d + 1] = Oacc[dn8][3];
|
||||
}
|
||||
|
||||
@@ -1,35 +1,6 @@
|
||||
#include "attn_prefill_split_q.cuh"
|
||||
#include "attn_dispatchers.cuh"
|
||||
#include "attn_entry_utils.cuh"
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
#include "attn_prefill_split_q_mma.cuh"
|
||||
#endif
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static void dispatch_prefill(AttentionParams<bf16>& p) {
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
constexpr int WARPS = 4, BR = 16;
|
||||
// KV tile: bigger tiles amortize the per-tile cp.async wait + barrier +
|
||||
// loop overhead over more tensor-core work (this kernel is latency-bound,
|
||||
// not compute/bandwidth-bound), so BC=32 wins ~6-8% over BC=16 for
|
||||
// D<=128. D=256 stays at 16: BC=32 double-buffered would need 64KB smem,
|
||||
// over the 48KB static cap.
|
||||
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
|
||||
dim3 grid((p.q_len + BR * WARPS - 1) / (BR * WARPS), p.q_head, p.batch);
|
||||
dim3 block(WARPS * 32, 1, 1);
|
||||
// Static shared memory — double-buffered K/V only (no sQ: Q goes direct
|
||||
// to registers). 2*BC*LD bf16 each for sK and sV → 4*BC*HEAD_DIM*2 bytes.
|
||||
// Occupancy is smem-capped: D=64→3 blocks/SM (16KB), D=128→1 (32KB),
|
||||
// D=256→1 (32KB, BC=16).
|
||||
attn_prefill_split_q_mma_kernel<HEAD_DIM, WARPS, BC><<<grid, block>>>(p);
|
||||
#else
|
||||
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
||||
dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
|
||||
dim3 block(G, ROWS, 1);
|
||||
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC><<<grid, block>>>(p);
|
||||
#endif
|
||||
}
|
||||
|
||||
torch::Tensor attn_prefill(
|
||||
torch::Tensor q,
|
||||
torch::Tensor k,
|
||||
|
||||
@@ -6,12 +6,9 @@
|
||||
using bf16 = __nv_bfloat16;
|
||||
|
||||
// v9: group-split register blocking. G threads cooperate on one query row,
|
||||
// each owning HEAD_DIM/G dims of qreg[]/acc[]. Small per-thread footprint keeps
|
||||
// occupancy high; the S dot product is reduced across the G-lane group with a
|
||||
// short shuffle chain (log2(G) shuffles) instead of a full 32-lane warp reduce.
|
||||
// Online (per-kv) softmax — cheap because acc[] is only HEAD_DIM/G long.
|
||||
// Templated on <HEAD_DIM, G, ROWS, P_BC>. Block = (G, ROWS). G power-of-two,
|
||||
// G*ROWS a multiple of 32 with groups warp-aligned.
|
||||
// each owning HEAD_DIM/G dims of qreg[]/acc[]. IsCausal and HasMask are
|
||||
// compile-time bools — the compiler eliminates dead branches.
|
||||
// Templated on <HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask>.
|
||||
|
||||
template <int G>
|
||||
__device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
|
||||
@@ -21,8 +18,7 @@ __device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
|
||||
return v;
|
||||
}
|
||||
|
||||
// load 8 contiguous bf16 from (16-byte aligned) smem as one float4, unpack to
|
||||
// 8 floats — cuts shared-load instructions 8x vs scalar bf16 loads.
|
||||
// load 8 contiguous bf16 from (16-byte aligned) smem as one float4
|
||||
__device__ __forceinline__ void ld8(const bf16* p, float* o) {
|
||||
float4 raw = *reinterpret_cast<const float4*>(p);
|
||||
const __nv_bfloat162* h = reinterpret_cast<const __nv_bfloat162*>(&raw);
|
||||
@@ -34,7 +30,7 @@ __device__ __forceinline__ void ld8(const bf16* p, float* o) {
|
||||
}
|
||||
}
|
||||
|
||||
template <int HEAD_DIM, int G, int ROWS, int P_BC>
|
||||
template <int HEAD_DIM, int G, int ROWS, int P_BC, bool IsCausal, bool HasMask>
|
||||
__global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
||||
constexpr int DPT = HEAD_DIM / G;
|
||||
|
||||
@@ -57,7 +53,7 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
||||
+ q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < DPT; i++)
|
||||
qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]) * p.scale;
|
||||
qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
|
||||
}
|
||||
|
||||
float m = -FLT_MAX, l = 0.0f;
|
||||
@@ -73,8 +69,6 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
||||
int tt = G * ROWS;
|
||||
int lid = row * G + gpos;
|
||||
|
||||
// per-group shuffle mask: only the G lanes of this row's group participate,
|
||||
// so causal masking (differing loop bounds across rows in a warp) is safe.
|
||||
int lane_in_warp = lid & 31;
|
||||
unsigned gmask = (G == 32) ? 0xFFFFFFFFu
|
||||
: (((1u << G) - 1u) << (lane_in_warp & ~(G - 1)));
|
||||
@@ -95,12 +89,14 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
||||
__syncthreads();
|
||||
|
||||
int lim = tlen;
|
||||
if (p.causal_offset >= 0 && q_row < p.q_len) {
|
||||
int ep = q_row + p.causal_offset + 1;
|
||||
if (kv0 >= ep)
|
||||
lim = 0;
|
||||
else if (kv0 + tlen > ep)
|
||||
lim = ep - kv0;
|
||||
if constexpr (IsCausal) {
|
||||
if (q_row < p.q_len) {
|
||||
int ep = q_row + p.causal_offset + 1;
|
||||
if (kv0 >= ep)
|
||||
lim = 0;
|
||||
else if (kv0 + tlen > ep)
|
||||
lim = ep - kv0;
|
||||
}
|
||||
}
|
||||
|
||||
int mask_row_base = mask_batch_base + q_row * p.mask_q_stride;
|
||||
@@ -115,11 +111,13 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
||||
for (int j = 0; j < 8; j++)
|
||||
part = fmaf(qreg[i + j], k8[j], part);
|
||||
}
|
||||
float dot = group_reduce_sum<G>(part, gmask);
|
||||
float dot = group_reduce_sum<G>(part, gmask) * p.scale;
|
||||
|
||||
int kv_idx = kv0 + s;
|
||||
if (p.use_mask && p.mask && !p.mask[mask_row_base + kv_idx])
|
||||
dot = -FLT_MAX;
|
||||
if constexpr (HasMask) {
|
||||
if (!p.mask[mask_row_base + kv_idx])
|
||||
dot = -FLT_MAX;
|
||||
}
|
||||
|
||||
float nm = fmaxf(m, dot);
|
||||
float al = __expf(m - nm);
|
||||
@@ -141,10 +139,9 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
||||
}
|
||||
|
||||
if (q_row < p.q_len) {
|
||||
// O: stride-based write
|
||||
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
||||
+ q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
|
||||
float rl = (l > 1e-10f) ? (1.0f / l) : 0.0f;
|
||||
float rl = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < DPT; i++)
|
||||
p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * rl);
|
||||
|
||||
@@ -4,121 +4,76 @@
|
||||
#include "attn_common.h"
|
||||
#include "attn_mma_utils.cuh"
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
|
||||
// Tensor-core prefill flash attention (raw mma.sync PTX).
|
||||
// One warp owns BR=16 query rows. S = Q@K^T and O = P@V run on bf16 tensor
|
||||
// cores via mma.sync.m16n8k16 (f32 accumulate). Q fragments are loaded once
|
||||
// straight from global into the mma A-operand layout (no smem staging) and
|
||||
// kept resident in registers across the tile loop. S, O, and the online-softmax
|
||||
// stats (m, l) also live in registers.
|
||||
// Shared memory is statically sized via template parameters — no dynamic
|
||||
// allocation. The mma fragment layout is used directly: the S accumulator
|
||||
// (f32) maps element-for-element onto the P matrix_a (bf16) operand, so
|
||||
// softmax needs no shuffle repack; row reductions fold across the 4-lane
|
||||
// thread group. Templated on <HEAD_DIM, WARPS, BC> with BC a multiple of 16.
|
||||
// cores via mma.sync.m16n8k16 (f32 accumulate).
|
||||
//
|
||||
// Software pipeline: K/V are double-buffered and loaded via cp.async one tile
|
||||
// ahead, so the next tile streams from global memory while the current tile's
|
||||
// tensor-core math runs — hiding load latency (long_scoreboard). A single
|
||||
// __syncthreads per tile both publishes the freshly loaded tile cross-warp and
|
||||
// (because it runs before the next prefetch) guards the buffer being refilled,
|
||||
// so no second barrier is needed. Predicated cp.async (cp_async_16_pred)
|
||||
// zero-fills rows past kv_len, unifying full and partial tiles on one path.
|
||||
// BC=32 (D<=128) amortizes the per-tile wait+barrier+loop overhead over more
|
||||
// tensor-core work — this kernel is latency-bound (low occupancy from high
|
||||
// register pressure), so fewer, larger tiles beat many tiny ones.
|
||||
// IsCausal and HasMask are compile-time bools — the compiler eliminates all
|
||||
// dead branches in the inner compute loop (FA2-style).
|
||||
//
|
||||
// Optimizations: load Q fragments directly from global in mma A-operand layout
|
||||
// (no sQ staging, no prologue barriers); post-multiply scale in float after
|
||||
// S=Q@K^T to avoid bf16 precision loss; packed bf16x2 output stores;
|
||||
// causal tile skipping (block-level prefetch bound + warp-level compute skip);
|
||||
// XOR swizzle (swiz_col) → eliminates ldmatrix bank conflicts without LD
|
||||
// padding (LD=HEAD_DIM).
|
||||
|
||||
template <int HEAD_DIM, int WARPS, int BC>
|
||||
// Traits = KernelTraits<HEAD_DIM, BC, WARPS=4, STAGES=2>.
|
||||
template <typename Traits, bool IsCausal, bool HasMask>
|
||||
__global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
|
||||
constexpr int BR = 16;
|
||||
constexpr int KD = HEAD_DIM / 16; // Q/K k-tiles
|
||||
constexpr int NC8 = BC / 8; // S n-tiles (N=8 each)
|
||||
constexpr int KT2 = BC / 16; // P k-tiles (K=16 each)
|
||||
constexpr int DN8 = HEAD_DIM / 8; // O n-tiles (N=8 each)
|
||||
constexpr int LD = HEAD_DIM; // XOR swizzle (swiz_col) handles bank conflicts
|
||||
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1); // chunk bits, stay within LD
|
||||
|
||||
const int warp = threadIdx.x / 32;
|
||||
const int lane = threadIdx.x % 32;
|
||||
const int gid = lane >> 2; // 0..7 → rows gid, gid+8
|
||||
const int gid = lane >> 2; // 0..7
|
||||
const int tid4 = lane & 3; // 0..3
|
||||
const int nthreads = WARPS * 32;
|
||||
|
||||
const int q_head = blockIdx.y;
|
||||
const int batch = blockIdx.z;
|
||||
const int kv_head = q_head / (p.q_head / p.kv_head);
|
||||
const int qrow0 = (blockIdx.x * WARPS + warp) * BR;
|
||||
const int qrow0 = (blockIdx.x * Traits::WARPS + warp) * Traits::BR;
|
||||
|
||||
// ---- Static shared memory: double-buffered K/V ----
|
||||
// K/V are double-buffered (STAGES=2): the next tile's cp.async load runs
|
||||
// while the current tile's tensor-core math executes, hiding global-load
|
||||
// latency (FA2-style software pipeline). No dynamic smem / carveout opt-in.
|
||||
constexpr int STAGES = 2;
|
||||
__shared__ __align__(16) bf16 sK[STAGES * BC * LD];
|
||||
__shared__ __align__(16) bf16 sV[STAGES * BC * LD];
|
||||
// Static shared memory: double-buffered K/V (no sQ — Q goes direct
|
||||
// to registers in mma A-operand layout).
|
||||
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
|
||||
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
|
||||
|
||||
// Load Q fragments straight from global into mma A-operand layout.
|
||||
// stride_row = p.q_stride_l for prefill (multi-q rows across q_len).
|
||||
// See attn_mma_utils.cuh for the shared template.
|
||||
const int q_base = batch * p.q_stride_b + q_head * p.q_stride_h;
|
||||
const int qra = qrow0 + gid;
|
||||
const int qrb = qrow0 + gid + 8;
|
||||
const bool va = qra < p.q_len, vb = qrb < p.q_len;
|
||||
unsigned Qa[KD][4];
|
||||
load_q_mma_frags<KD>(p.q + q_base, p.q_stride_l, p.q_stride_d,
|
||||
qra, qrb, va, vb, tid4, Qa);
|
||||
unsigned Qa[Traits::KD][4];
|
||||
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_l, p.q_stride_d,
|
||||
qra, qrb, va, vb, tid4, Qa);
|
||||
|
||||
float Oacc[DN8][4];
|
||||
#pragma unroll
|
||||
for (int j = 0; j < DN8; j++)
|
||||
float Oacc[Traits::DN8][4];
|
||||
#pragma unroll
|
||||
for (int j = 0; j < Traits::DN8; j++)
|
||||
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
||||
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
||||
|
||||
// KV: stride-based base
|
||||
const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||
const int tiles = (p.kv_len + BC - 1) / BC;
|
||||
const int qr0 = qrow0 + gid; // row for c0/c1
|
||||
const int qr1 = qrow0 + gid + 8; // row for c2/c3
|
||||
const int tiles = (p.kv_len + Traits::BC - 1) / Traits::BC;
|
||||
const int qr0 = qrow0 + gid;
|
||||
const int qr1 = qrow0 + gid + 8;
|
||||
|
||||
// Causal tile-skip bounds (no-op when causal_offset < 0)
|
||||
const int use_skip = (p.causal_offset >= 0) ? 1 : 0;
|
||||
const int max_kv = qrow0 + BR - 1 + p.causal_offset;
|
||||
// Causal tile-skip bounds (dead code when IsCausal == false)
|
||||
const int max_kv = qrow0 + Traits::BR - 1 + p.causal_offset;
|
||||
const int block_max_kv =
|
||||
blockIdx.x * WARPS * BR + WARPS * BR - 1 + p.causal_offset;
|
||||
const int has_mask = p.use_mask && p.mask;
|
||||
blockIdx.x * Traits::WARPS * Traits::BR + Traits::WARPS * Traits::BR - 1
|
||||
+ p.causal_offset;
|
||||
|
||||
// Last active tile: block-level causal bound (all warps in the block share
|
||||
// the K/V load, so the prefetch range is the block max, not per-warp).
|
||||
int t_end = tiles - 1;
|
||||
if (use_skip) {
|
||||
int bt = block_max_kv / BC;
|
||||
if constexpr (IsCausal) {
|
||||
int bt = block_max_kv / Traits::BC;
|
||||
if (bt < t_end) t_end = bt;
|
||||
}
|
||||
|
||||
constexpr int VEC = 8; // bf16 per cp.async unit (16 bytes)
|
||||
constexpr int TOTAL = BC * HEAD_DIM;
|
||||
|
||||
// ---- Load tile lambda: predicated cp.async ----
|
||||
// Issue cp.async loads for tile `ti` into shared buffer `buf`. Predicated
|
||||
// loads zero-fill rows past kv_len, so partial tiles need no scalar path.
|
||||
auto load_tile = [&](int ti, int buf) {
|
||||
int kv0 = ti * BC;
|
||||
bf16* dK = sK + buf * BC * LD;
|
||||
bf16* dV = sV + buf * BC * LD;
|
||||
#pragma unroll
|
||||
for (int i = threadIdx.x * VEC; i < TOTAL; i += nthreads * VEC) {
|
||||
int r = i / HEAD_DIM, d = i % HEAD_DIM;
|
||||
int kv0 = ti * Traits::BC;
|
||||
bf16* dK = sK + buf * Traits::BC * Traits::LD;
|
||||
bf16* dV = sV + buf * Traits::BC * Traits::LD;
|
||||
#pragma unroll
|
||||
for (int i = threadIdx.x * 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;
|
||||
int off = r * LD + swiz_col(d, r, SWIZ_MASK);
|
||||
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
|
||||
int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d;
|
||||
cp_async_16_pred(&dK[off], &p.k[g_off], valid);
|
||||
cp_async_16_pred(&dV[off], &p.v[g_off], valid);
|
||||
@@ -132,65 +87,60 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
|
||||
for (int ti = 0; ti <= t_end; ti++) {
|
||||
int buf = ti & 1;
|
||||
|
||||
// Wait for the current tile's async copies, then a single barrier: it
|
||||
// both publishes this tile's data cross-warp AND guarantees the prior
|
||||
// compute on the buffer we are about to refill has finished. Issuing
|
||||
// the next tile's load *after* this barrier lets one barrier cover both
|
||||
// hazards (vs two), while the load still overlaps this tile's math.
|
||||
// Wait for current tile, then publish cross-warp + guard buffer reuse.
|
||||
cp_async_wait_group<0>();
|
||||
__syncthreads();
|
||||
if (ti < t_end) load_tile(ti + 1, (ti + 1) & 1);
|
||||
|
||||
const bf16* bK = sK + buf * BC * LD;
|
||||
const bf16* bV = sV + buf * BC * LD;
|
||||
int kv0 = ti * BC;
|
||||
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
|
||||
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
|
||||
int kv0 = ti * Traits::BC;
|
||||
|
||||
// Warp-level causal skip
|
||||
if (!use_skip || kv0 <= max_kv) {
|
||||
// Warp-level causal skip (dead branch eliminated when IsCausal == false)
|
||||
if (!IsCausal || kv0 <= max_kv) {
|
||||
|
||||
// S = Q @ K^T + scale + online softmax + O += P @ V
|
||||
float Sacc[NC8][4];
|
||||
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
|
||||
float Sacc[Traits::NC8][4];
|
||||
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
|
||||
|
||||
// post-multiply scale in float (no bf16 precision loss from pre-scaling Q)
|
||||
#pragma unroll
|
||||
for (int n8 = 0; n8 < NC8; n8++)
|
||||
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||
// Post-multiply scale in float (no bf16 precision loss)
|
||||
#pragma unroll
|
||||
for (int n8 = 0; n8 < Traits::NC8; n8++)
|
||||
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||
|
||||
int maxc0 = (p.causal_offset >= 0) ? min(p.kv_len, qr0 + p.causal_offset + 1)
|
||||
: p.kv_len;
|
||||
int maxc1 = (p.causal_offset >= 0) ? min(p.kv_len, qr1 + p.causal_offset + 1)
|
||||
: p.kv_len;
|
||||
mma_softmax_tile<NC8, DN8>(kv0, maxc0, maxc1,
|
||||
qr0, qr1,
|
||||
p.mask_b_stride, p.mask_q_stride,
|
||||
batch,
|
||||
p.mask, has_mask,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
int maxc0 = IsCausal ? min(p.kv_len, qr0 + p.causal_offset + 1)
|
||||
: p.kv_len;
|
||||
int maxc1 = IsCausal ? min(p.kv_len, qr1 + p.causal_offset + 1)
|
||||
: p.kv_len;
|
||||
mma_softmax_tile<Traits, HasMask>(kv0, maxc0, maxc1,
|
||||
qr0, qr1,
|
||||
p.mask_b_stride, p.mask_q_stride,
|
||||
batch,
|
||||
p.mask,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
|
||||
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc);
|
||||
} // if active (warp-level causal skip)
|
||||
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
||||
}
|
||||
}
|
||||
|
||||
// ---- write output ---- (packed bf16x2 stores: one 32-bit STG per pair,
|
||||
// halves store count and removes the uncoalesced scalar-store penalty)
|
||||
// ---- write output: packed bf16x2 stores ----
|
||||
float rl0 = (l0 > 1e-20f) ? (1.0f / l0) : 0.0f;
|
||||
float rl1 = (l1 > 1e-20f) ? (1.0f / l1) : 0.0f;
|
||||
// O: stride-based write
|
||||
const int o_base = batch * p.q_stride_b + q_head * p.q_stride_h;
|
||||
#pragma unroll
|
||||
for (int dn8 = 0; dn8 < DN8; dn8++) {
|
||||
#pragma unroll
|
||||
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
|
||||
int d = dn8 * 8 + 2 * tid4;
|
||||
if (qr0 < p.q_len) {
|
||||
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][0] * rl0,
|
||||
Oacc[dn8][1] * rl0);
|
||||
*reinterpret_cast<__nv_bfloat162*>(&p.o[o_base + qr0 * p.q_stride_l + d * p.q_stride_d]) = v;
|
||||
Oacc[dn8][1] * rl0);
|
||||
*reinterpret_cast<__nv_bfloat162*>(
|
||||
&p.o[o_base + qr0 * p.q_stride_l + d * p.q_stride_d]) = v;
|
||||
}
|
||||
if (qr1 < p.q_len) {
|
||||
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1,
|
||||
Oacc[dn8][3] * rl1);
|
||||
*reinterpret_cast<__nv_bfloat162*>(&p.o[o_base + qr1 * p.q_stride_l + d * p.q_stride_d]) = v;
|
||||
Oacc[dn8][3] * rl1);
|
||||
*reinterpret_cast<__nv_bfloat162*>(
|
||||
&p.o[o_base + qr1 * p.q_stride_l + d * p.q_stride_d]) = v;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
#pragma once
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
|
||||
static constexpr int MAX_SPLITS = 32;
|
||||
|
||||
__device__ inline float warp_reduce_sum(float val) {
|
||||
#pragma unroll
|
||||
for (int offset = 16; offset > 0; offset >>= 1)
|
||||
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
|
||||
return val;
|
||||
}
|
||||
+102
-118
@@ -1,73 +1,33 @@
|
||||
/*
|
||||
Pure-C test:
|
||||
Pure-C test — uses shared dispatcher.
|
||||
nvcc -I csrc -arch=sm_89 -O3 \
|
||||
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
|
||||
csrc/tests/attn_decode_test.cu -o test && ./test
|
||||
*/
|
||||
|
||||
#include "test_utils.cuh"
|
||||
#include "../kernels/attn_decode_split_kv.cuh"
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
#include "../kernels/attn_decode_split_kv_mma.cuh"
|
||||
#endif
|
||||
#include "../kernels/attn_dispatchers.cuh"
|
||||
|
||||
// Split-K scratch (torch-free): the production launcher allocates these from
|
||||
// torch; here we pass pre-allocated device buffers so the bench loop doesn't
|
||||
// pay a cudaMalloc per iteration. Size for the maximum split count (32).
|
||||
// Split-K scratch (torch-free)
|
||||
struct DecodeScratch {
|
||||
float* o_part = nullptr;
|
||||
float* ml_part = nullptr;
|
||||
};
|
||||
|
||||
// Launch the production decode path (tensor-core head-packing MMA on sm_80+,
|
||||
// scalar fallback otherwise), mirroring dispatch_decode() in attn_decode.cu.
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
static bool decode_use_mma(const AttentionParams<bf16>& p) {
|
||||
int G = p.q_head / p.kv_head;
|
||||
return !p.use_mask && G > 1 && G <= 16;
|
||||
static void setup_scratch(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
||||
int max_splits = 32;
|
||||
cudaMalloc(&sc.o_part, (size_t)p.batch * p.q_head * max_splits * p.head_dim * sizeof(float));
|
||||
cudaMalloc(&sc.ml_part, (size_t)p.batch * p.q_head * max_splits * 2 * sizeof(float));
|
||||
}
|
||||
|
||||
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
|
||||
static void launch_mma_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
||||
int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
||||
p.o_part = sc.o_part;
|
||||
p.ml_part = sc.ml_part;
|
||||
|
||||
attn_decode_split_kv_mma_kernel<HEAD_DIM, BC, STAGES>
|
||||
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
|
||||
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
#endif
|
||||
|
||||
static void launch_scalar_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
||||
int gs = p.q_head / p.kv_head;
|
||||
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||
p.o_part = sc.o_part;
|
||||
p.ml_part = sc.ml_part;
|
||||
|
||||
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
||||
attn_decode_split_kv_kernel<<<dim3(p.batch * p.kv_head, 1, p.num_splits), dim3(32, gs), smem>>>(p);
|
||||
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static void dispatch_decode_t(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
if (decode_use_mma(p)) { launch_mma_decode<HEAD_DIM, 32>(p, sc); return; }
|
||||
#endif
|
||||
launch_scalar_decode(p, sc);
|
||||
}
|
||||
|
||||
static void dispatch_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
||||
dispatch_by_head_dim(p.head_dim, [&]<int D>() { dispatch_decode_t<D>(p, sc); });
|
||||
static void free_scratch(DecodeScratch& sc) {
|
||||
cudaFree(sc.o_part); cudaFree(sc.ml_part);
|
||||
}
|
||||
|
||||
// Warmed-up, CUDA-event timed sweep over the production decode MMA path.
|
||||
static void bench() {
|
||||
const int cfgs[][5] = {
|
||||
{1, 32, 4, 512, 128}, // B, Hq, Hk, kv_len, D
|
||||
{1, 32, 4, 512, 128},
|
||||
{1, 32, 4, 1024, 128},
|
||||
{1, 32, 4, 2048, 128},
|
||||
{1, 32, 4, 4096, 128},
|
||||
@@ -104,10 +64,10 @@ static void bench() {
|
||||
p.q = dQ; p.k = dK; p.v = dV; p.mask = nullptr; p.o = dO;
|
||||
|
||||
DecodeScratch sc;
|
||||
cudaMalloc(&sc.o_part, (size_t)B*Hq*32*D*sizeof(float));
|
||||
cudaMalloc(&sc.ml_part, (size_t)B*Hq*32*2*sizeof(float));
|
||||
setup_scratch(p, sc);
|
||||
p.o_part = sc.o_part; p.ml_part = sc.ml_part;
|
||||
|
||||
auto launch = [&]() { dispatch_decode(p, sc); };
|
||||
auto launch = [&]() { dispatch_by_head_dim(D, [&]<int H>() { dispatch_decode<H>(p); }); };
|
||||
double flops = 4.0 * B * Hq * (double)sl * D;
|
||||
double bytes = 2.0 * (2.0 * nKV * sizeof(bf16));
|
||||
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops, bytes);
|
||||
@@ -119,81 +79,105 @@ static void bench() {
|
||||
print_bench_row(cfg, r);
|
||||
|
||||
cudaFree(dQ); cudaFree(dK); cudaFree(dV); cudaFree(dO);
|
||||
cudaFree(sc.o_part); cudaFree(sc.ml_part);
|
||||
free_scratch(sc);
|
||||
}
|
||||
}
|
||||
|
||||
static int run_test(int B, int Hq, int Hk, int sl, int D, int causal) {
|
||||
int gs = Hq / Hk;
|
||||
printf("=== B=%d Hq=%d Hk=%d seq=%d D=%d gs=%d causal=%d ===\n",
|
||||
B,Hq,Hk,sl,D,gs,causal);
|
||||
|
||||
size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D;
|
||||
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
||||
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
||||
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
||||
|
||||
bool* hMask=new bool[B*sl];
|
||||
for (int i=0;i<B*sl;i++) hMask[i]=true;
|
||||
|
||||
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||
bool* dMask;
|
||||
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||
cudaMalloc(&dMask,B*sl);
|
||||
|
||||
tmp=new bf16[max(nQ,nKV)];
|
||||
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
||||
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
||||
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
||||
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
cudaMemcpy(dMask,hMask,B*sl,cudaMemcpyHostToDevice);
|
||||
|
||||
AttentionParams<bf16> p;
|
||||
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D;
|
||||
p.use_mask=0; p.causal_offset=causal?0:-1;
|
||||
p.scale=1.0f/sqrtf((float)D);
|
||||
set_default_strides(p);
|
||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||
|
||||
DecodeScratch sc;
|
||||
setup_scratch(p, sc);
|
||||
p.o_part = sc.o_part; p.ml_part = sc.ml_part;
|
||||
|
||||
double t0=now_ms();
|
||||
dispatch_by_head_dim(D, [&]<int H>() { dispatch_decode<H>(p); });
|
||||
cudaDeviceSynchronize();
|
||||
double kms=now_ms()-t0;
|
||||
cudaError_t err=cudaGetLastError();
|
||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
||||
|
||||
bf16* hOut=new bf16[nQ];
|
||||
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
||||
|
||||
float* ref=new float[nQ];
|
||||
cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, causal ? 0 : -1);
|
||||
|
||||
float max_abs_err=0, max_rel_err=0;
|
||||
for (size_t i=0;i<nQ;i++){
|
||||
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
||||
if(err>max_abs_err) max_abs_err=err;
|
||||
float rel=err/fmaxf(fabsf(ref[i]), 1e-8f);
|
||||
if(rel>max_rel_err) max_rel_err=rel;
|
||||
}
|
||||
const float atol=0.01f, rtol=0.01f;
|
||||
bool pass=true;
|
||||
for (size_t i=0;i<nQ;i++){
|
||||
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
||||
if (err > atol + rtol * fabsf(ref[i])) { pass=false; break; }
|
||||
}
|
||||
printf("kernel: %.3f ms max_abs_err: %.6e max_rel_err: %.6e %s\n\n",
|
||||
kms, max_abs_err, max_rel_err, pass?"PASS":"FAIL");
|
||||
|
||||
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask);
|
||||
free_scratch(sc);
|
||||
delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp;
|
||||
|
||||
return pass ? 0 : 1;
|
||||
}
|
||||
|
||||
int main() {
|
||||
const int configs[][5] = {
|
||||
{1, 2, 1, 64, 32}, // B,Hq,Hk,seq_len,D
|
||||
{1, 32, 4, 512, 128},
|
||||
{1, 32, 4, 1024, 128},
|
||||
const int configs[][6] = {
|
||||
{1, 2, 1, 64, 32, 0},
|
||||
{1, 32, 4, 512, 128, 0},
|
||||
{1, 32, 4, 1024, 128, 0},
|
||||
{1, 32, 4, 512, 128, 1},
|
||||
};
|
||||
int n_cfgs = sizeof(configs) / sizeof(configs[0]);
|
||||
int fail = 0;
|
||||
|
||||
for (int ci = 0; ci < n_cfgs; ci++) {
|
||||
int B = configs[ci][0], Hq = configs[ci][1], Hk = configs[ci][2];
|
||||
int sl = configs[ci][3], D = configs[ci][4], gs = Hq / Hk;
|
||||
printf("=== B=%d Hq=%d Hk=%d seq=%d D=%d gs=%d ===\n", B,Hq,Hk,sl,D,gs);
|
||||
int sl = configs[ci][3], D = configs[ci][4], causal = configs[ci][5];
|
||||
fail += run_test(B, Hq, Hk, sl, D, causal);
|
||||
if (fail) break;
|
||||
}
|
||||
|
||||
size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D;
|
||||
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
||||
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
||||
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
||||
|
||||
bool* hMask=new bool[B*sl];
|
||||
for (int i=0;i<B*sl;i++) hMask[i]=true;
|
||||
|
||||
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||
bool* dMask;
|
||||
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||
cudaMalloc(&dMask,B*sl);
|
||||
|
||||
tmp=new bf16[max(nQ,nKV)];
|
||||
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
||||
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
||||
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
||||
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
cudaMemcpy(dMask,hMask,B*sl,cudaMemcpyHostToDevice);
|
||||
|
||||
AttentionParams<bf16> p;
|
||||
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D;
|
||||
p.use_mask=0; p.causal_offset=-1;
|
||||
p.scale=1.0f/sqrtf((float)D);
|
||||
set_default_strides(p);
|
||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||
|
||||
// Split-K scratch (max 32 splits), sized for the production MMA path.
|
||||
DecodeScratch sc;
|
||||
cudaMalloc(&sc.o_part, (size_t)B*Hq*32*D*sizeof(float));
|
||||
cudaMalloc(&sc.ml_part, (size_t)B*Hq*32*2*sizeof(float));
|
||||
|
||||
double t0=now_ms();
|
||||
dispatch_decode(p, sc);
|
||||
cudaDeviceSynchronize();
|
||||
double kms=now_ms()-t0;
|
||||
cudaError_t err=cudaGetLastError();
|
||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
||||
|
||||
bf16* hOut=new bf16[nQ];
|
||||
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
||||
|
||||
float* ref=new float[nQ];
|
||||
cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, -1);
|
||||
|
||||
float max_err=0;
|
||||
for (size_t i=0;i<nQ;i++){
|
||||
float d=fabsf(bf2f(hOut[i])-ref[i]);
|
||||
if(d>max_err) max_err=d;
|
||||
}
|
||||
printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err);
|
||||
|
||||
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask);
|
||||
cudaFree(sc.o_part);cudaFree(sc.ml_part);
|
||||
delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp;
|
||||
if (fail) {
|
||||
printf("FAILED\n");
|
||||
return fail;
|
||||
}
|
||||
printf("All tests passed!\n");
|
||||
bench();
|
||||
|
||||
@@ -5,12 +5,8 @@
|
||||
|
||||
#include <cstring>
|
||||
#include "test_utils.cuh"
|
||||
#include "../kernels/attn_paged_decode_split_kv.cuh"
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
#include "../kernels/attn_paged_decode_split_kv_mma.cuh"
|
||||
#endif
|
||||
#include "../kernels/attn_dispatchers.cuh"
|
||||
|
||||
// Copy contiguous K/V from page pool (reference gather)
|
||||
static void gather_kv_cpu(
|
||||
const bf16* h_k_pool, const bf16* h_v_pool,
|
||||
const int64_t* h_pt, int B, int Hkv, int kv_len,
|
||||
@@ -28,7 +24,8 @@ static void gather_kv_cpu(
|
||||
size_t src_base = (size_t)phys * page_stride
|
||||
+ (size_t)pg_off * Hkv * head_dim
|
||||
+ h * head_dim;
|
||||
size_t dst_base = ((size_t)b * Hkv + h) * kv_len * head_dim + (size_t)pos * head_dim;
|
||||
size_t dst_base = ((size_t)b * Hkv + h) * kv_len * head_dim
|
||||
+ (size_t)pos * head_dim;
|
||||
memcpy(h_k + dst_base, h_k_pool + src_base, head_dim * sizeof(bf16));
|
||||
memcpy(h_v + dst_base, h_v_pool + src_base, head_dim * sizeof(bf16));
|
||||
}
|
||||
@@ -37,54 +34,29 @@ static void gather_kv_cpu(
|
||||
}
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static void launch_paged_decode(PagedAttentionParams<bf16, float>& p) {
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
int G_check = p.q_head / p.kv_head;
|
||||
bool use_mma = !p.use_mask && G_check >= 1 && G_check <= 16 && p.page_size >= 32;
|
||||
if (use_mma) {
|
||||
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
|
||||
int tiles_total = (p.kv_len + 32 - 1) / 32;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
||||
paged_attn_decode_split_kv_mma_kernel<HEAD_DIM, 32, STAGES>
|
||||
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
|
||||
} else
|
||||
#endif
|
||||
{
|
||||
int group_sz = p.q_head / p.kv_head;
|
||||
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);
|
||||
paged_attn_decode_split_kv_kernel<<<
|
||||
dim3(p.batch * p.kv_head, 1, p.num_splits),
|
||||
dim3(32, group_sz), smem>>>(p);
|
||||
}
|
||||
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed) {
|
||||
printf("B=%d Hq=%d Hkv=%d kv_len=%d page_sz=%d head_dim=%d ... ", B, Hq, Hkv, kv_len, page_size, HEAD_DIM);
|
||||
static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int causal, int seed) {
|
||||
printf("B=%d Hq=%d Hkv=%d kv_len=%d page_sz=%d head_dim=%d causal=%d ... ",
|
||||
B, Hq, Hkv, kv_len, page_size, HEAD_DIM, causal);
|
||||
fflush(stdout);
|
||||
|
||||
int max_pages = (kv_len + page_size - 1) / page_size;
|
||||
int n_phys_pages = B * max_pages;
|
||||
int max_splits = 32;
|
||||
|
||||
size_t sz_q = (size_t)B * Hq * 1 * HEAD_DIM * sizeof(bf16);
|
||||
size_t sz_o = sz_q;
|
||||
size_t sz_kv = (size_t)n_phys_pages * page_size * Hkv * HEAD_DIM * sizeof(bf16);
|
||||
size_t sz_pt = (size_t)B * max_pages * sizeof(int64_t);
|
||||
int max_splits = 32;
|
||||
size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float);
|
||||
size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float);
|
||||
|
||||
bf16 *d_q, *d_o_paged, *d_o_ref;
|
||||
bf16 *d_q, *d_o_paged;
|
||||
bf16 *d_k_pool, *d_v_pool;
|
||||
int64_t* d_pt;
|
||||
float *d_op, *d_ml;
|
||||
|
||||
cudaMalloc(&d_q, sz_q);
|
||||
cudaMalloc(&d_o_paged, sz_o);
|
||||
cudaMalloc(&d_o_ref, sz_o);
|
||||
cudaMalloc(&d_k_pool, sz_kv);
|
||||
cudaMalloc(&d_v_pool, sz_kv);
|
||||
cudaMalloc(&d_pt, sz_pt);
|
||||
@@ -107,7 +79,8 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed)
|
||||
for (int h = 0; h < Hkv; h++) {
|
||||
for (int d = 0; d < HEAD_DIM; d++) {
|
||||
float v = sinf((float)(pg * 7919 + off * 1049 + h * 331 + d));
|
||||
size_t idx = (size_t)pg * ps + (size_t)off * Hkv * HEAD_DIM + h * HEAD_DIM + d;
|
||||
size_t idx = (size_t)pg * ps + (size_t)off * Hkv * HEAD_DIM
|
||||
+ h * HEAD_DIM + d;
|
||||
h_k_pool[idx] = __float2bfloat16(v);
|
||||
h_v_pool[idx] = __float2bfloat16(v * 0.3f);
|
||||
}
|
||||
@@ -138,22 +111,22 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed)
|
||||
}
|
||||
|
||||
float* h_o_ref = (float*)calloc(B * Hq * HEAD_DIM, sizeof(float));
|
||||
cpu_attention_ref(h_q_f, h_k_f, h_v_f, nullptr, h_o_ref, B, Hq, Hkv, 1, kv_len, HEAD_DIM, -1);
|
||||
cpu_attention_ref(h_q_f, h_k_f, h_v_f, nullptr, h_o_ref, B, Hq, Hkv,
|
||||
1, kv_len, HEAD_DIM, causal ? 0 : -1);
|
||||
|
||||
float scale_val = 1.0f / sqrtf((float)HEAD_DIM);
|
||||
PagedAttentionParams<bf16, float> p;
|
||||
PagedAttentionParams<bf16> p;
|
||||
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.q_len = 1;
|
||||
p.kv_len = kv_len; p.head_dim = HEAD_DIM;
|
||||
p.use_mask = 0; p.causal_offset = -1;
|
||||
p.use_mask = 0; p.causal_offset = causal ? 0 : -1;
|
||||
set_default_paged_strides(p);
|
||||
p.num_splits = 1; p.scale = scale_val;
|
||||
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||
p.page_size = page_size; p.max_pages = max_pages;
|
||||
p.page_table = d_pt;
|
||||
p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||
p.q = d_q; p.mask = nullptr; p.o = d_o_paged;
|
||||
p.o_part = d_op; p.ml_part = d_ml;
|
||||
|
||||
launch_paged_decode<HEAD_DIM>(p);
|
||||
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(p); });
|
||||
cudaDeviceSynchronize();
|
||||
|
||||
bf16* h_o_bf16 = (bf16*)malloc(sz_o);
|
||||
@@ -162,23 +135,30 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed)
|
||||
for (int i = 0; i < B * Hq * HEAD_DIM; i++)
|
||||
h_o_paged[i] = __bfloat162float(h_o_bf16[i]);
|
||||
|
||||
float max_err = 0.0f;
|
||||
float max_abs_err = 0.0f, max_rel_err = 0.0f;
|
||||
int bad_idx = -1;
|
||||
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
|
||||
float e = fabsf(h_o_paged[i] - h_o_ref[i]);
|
||||
if (e > max_err) { max_err = e; bad_idx = i; }
|
||||
if (e > max_abs_err) { max_abs_err = e; bad_idx = i; }
|
||||
float rel = e / fmaxf(fabsf(h_o_ref[i]), 1e-8f);
|
||||
if (rel > max_rel_err) max_rel_err = rel;
|
||||
}
|
||||
|
||||
bool pass = max_err < 0.02f;
|
||||
const float atol = 0.01f, rtol = 0.01f;
|
||||
bool pass = true;
|
||||
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
|
||||
float e = fabsf(h_o_paged[i] - h_o_ref[i]);
|
||||
if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; }
|
||||
}
|
||||
|
||||
if (pass) {
|
||||
printf("PASS (max_abs_err=%.4e)\n", max_err);
|
||||
printf("PASS (max_abs_err=%.4e max_rel_err=%.4e)\n", max_abs_err, max_rel_err);
|
||||
} else {
|
||||
int b = bad_idx / (Hq * HEAD_DIM);
|
||||
int h = (bad_idx / HEAD_DIM) % Hq;
|
||||
int d = bad_idx % HEAD_DIM;
|
||||
printf("FAIL (max_abs_err=%.4e at [%d,%d,%d]: ref=%.4f got=%.4f)\n",
|
||||
max_err, b, h, d, h_o_ref[bad_idx], h_o_paged[bad_idx]);
|
||||
printf("FAIL (max_abs_err=%.4e max_rel_err=%.4e at [%d,%d,%d]: ref=%.4f got=%.4f)\n",
|
||||
max_abs_err, max_rel_err, b, h, d, h_o_ref[bad_idx], h_o_paged[bad_idx]);
|
||||
printf(" ref[0..7]:");
|
||||
for (int i = 0; i < 8 && i < HEAD_DIM; i++)
|
||||
printf(" %.4f", h_o_ref[i]);
|
||||
@@ -192,7 +172,7 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed)
|
||||
free(h_k_cont); free(h_v_cont);
|
||||
free(h_q_f); free(h_k_f); free(h_v_f);
|
||||
free(h_o_ref); free(h_o_bf16); free(h_o_paged);
|
||||
cudaFree(d_q); cudaFree(d_o_paged); cudaFree(d_o_ref);
|
||||
cudaFree(d_q); cudaFree(d_o_paged);
|
||||
cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_pt);
|
||||
cudaFree(d_op); cudaFree(d_ml);
|
||||
|
||||
@@ -201,48 +181,43 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed)
|
||||
|
||||
struct TestCase {
|
||||
int head_dim;
|
||||
int B, Hq, Hkv, kv_len, page_size, seed;
|
||||
int B, Hq, Hkv, kv_len, page_size, causal, seed;
|
||||
};
|
||||
|
||||
static const TestCase TESTS[] = {
|
||||
{128, 1, 1, 1, 8, 128, 1},
|
||||
{128, 1, 4, 4, 128, 128, 2},
|
||||
{128, 2, 4, 4, 256, 128, 3},
|
||||
{128, 1, 4, 1, 64, 64, 4},
|
||||
{128, 1, 8, 2, 64, 128, 5},
|
||||
{128, 2, 16, 4, 128, 128, 6},
|
||||
{64, 1, 4, 2, 32, 128, 7},
|
||||
{256, 1, 2, 1, 16, 128, 8},
|
||||
{32, 1, 4, 2, 32, 64, 9},
|
||||
{128, 3, 8, 2, 256, 128, 10},
|
||||
{128, 2, 32, 8, 512, 128, 11},
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
{128, 1, 16, 2, 256, 128, 12},
|
||||
{128, 2, 32, 4, 512, 128, 13},
|
||||
#endif
|
||||
{128, 1, 1, 1, 8, 128, 0, 1},
|
||||
{128, 1, 4, 4, 128, 128, 0, 2},
|
||||
{128, 2, 4, 4, 256, 128, 0, 3},
|
||||
{128, 1, 4, 1, 64, 64, 0, 4},
|
||||
{128, 1, 8, 2, 64, 128, 0, 5},
|
||||
{128, 2, 16, 4, 128, 128, 0, 6},
|
||||
{64, 1, 4, 2, 32, 128, 0, 7},
|
||||
{256, 1, 2, 1, 16, 128, 0, 8},
|
||||
{32, 1, 4, 2, 32, 64, 0, 9},
|
||||
{128, 3, 8, 2, 256, 128, 0, 10},
|
||||
{128, 2, 32, 8, 512, 128, 0, 11},
|
||||
{128, 1, 16, 2, 256, 128, 0, 12},
|
||||
{128, 2, 32, 4, 512, 128, 0, 13},
|
||||
{128, 2, 8, 2, 128, 128, 1, 14}, // causal
|
||||
};
|
||||
|
||||
static int dispatch_test(const TestCase& tc) {
|
||||
bool matched = false;
|
||||
int r = 0;
|
||||
dispatch_by_head_dim(tc.head_dim, [&]<int D>() {
|
||||
matched = true;
|
||||
r = run_test<D>(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.seed);
|
||||
r = run_test<D>(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.causal, tc.seed);
|
||||
});
|
||||
return matched ? r : 1;
|
||||
return r;
|
||||
}
|
||||
|
||||
// Warmed-up, CUDA-event timed sweep over paged decode configs.
|
||||
// Bytes = K + V read through page table (B*Hk*kv*D each), bf16.
|
||||
template <int HEAD_DIM>
|
||||
static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) {
|
||||
int max_pages = (kv_len + page_size - 1) / page_size;
|
||||
int n_phys_pages = B * max_pages;
|
||||
int max_splits = 32;
|
||||
|
||||
size_t sz_q = (size_t)B * Hq * 1 * HEAD_DIM * sizeof(bf16);
|
||||
size_t sz_kv = (size_t)n_phys_pages * page_size * Hkv * HEAD_DIM * sizeof(bf16);
|
||||
size_t sz_pt = (size_t)B * max_pages * sizeof(int64_t);
|
||||
int max_splits = 32;
|
||||
size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float);
|
||||
size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float);
|
||||
|
||||
@@ -269,13 +244,12 @@ static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) {
|
||||
cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice);
|
||||
free(h_pt);
|
||||
|
||||
float scale_val = 1.0f / sqrtf((float)HEAD_DIM);
|
||||
PagedAttentionParams<bf16, float> pa;
|
||||
PagedAttentionParams<bf16> pa;
|
||||
pa.batch = B; pa.q_head = Hq; pa.kv_head = Hkv; pa.q_len = 1;
|
||||
pa.kv_len = kv_len; pa.head_dim = HEAD_DIM;
|
||||
pa.use_mask = 0; pa.causal_offset = -1;
|
||||
set_default_paged_strides(pa);
|
||||
pa.num_splits = 1; pa.scale = scale_val;
|
||||
pa.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||
pa.page_size = page_size; pa.max_pages = max_pages;
|
||||
pa.page_table = d_pt;
|
||||
pa.k_cache = d_k_pool; pa.v_cache = d_v_pool;
|
||||
@@ -283,7 +257,9 @@ static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) {
|
||||
pa.o_part = d_op; pa.ml_part = d_ml;
|
||||
|
||||
const int WARMUP = 10, ITERS = 100;
|
||||
auto launch = [&]() { launch_paged_decode<HEAD_DIM>(pa); };
|
||||
auto launch = [&]() {
|
||||
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(pa); });
|
||||
};
|
||||
double flops = 4.0 * B * Hq * (double)kv_len * HEAD_DIM;
|
||||
size_t nKV = (size_t)B * Hkv * kv_len * HEAD_DIM;
|
||||
double bytes = 2.0 * (2.0 * nKV * sizeof(bf16));
|
||||
|
||||
@@ -1,45 +1,14 @@
|
||||
/*
|
||||
Pure-C test:
|
||||
Pure-C test — uses shared dispatcher.
|
||||
nvcc -I csrc -arch=sm_89 -O3 \
|
||||
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
|
||||
csrc/tests/attn_prefill_test.cu -o test && ./test
|
||||
*/
|
||||
|
||||
#include "test_utils.cuh"
|
||||
#include "../kernels/attn_prefill_split_q.cuh"
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
#include "../kernels/attn_prefill_split_q_mma.cuh"
|
||||
#endif
|
||||
|
||||
// Launch the production prefill path (tensor-core MMA on sm_80+, else the
|
||||
// scalar fallback), mirroring dispatch_prefill() in attn_prefill.cu.
|
||||
template <int HEAD_DIM>
|
||||
static void launch_prefill(AttentionParams<bf16>& p) {
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
constexpr int WARPS = 4, BR = 16;
|
||||
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
|
||||
dim3 grid((p.q_len + BR * WARPS - 1) / (BR * WARPS), p.q_head, p.batch);
|
||||
dim3 block(WARPS * 32, 1, 1);
|
||||
attn_prefill_split_q_mma_kernel<HEAD_DIM, WARPS, BC><<<grid, block>>>(p);
|
||||
#else
|
||||
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
||||
dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
|
||||
dim3 block(G, ROWS, 1);
|
||||
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC><<<grid, block>>>(p);
|
||||
#endif
|
||||
}
|
||||
|
||||
static void dispatch_prefill(AttentionParams<bf16>& p) {
|
||||
switch (p.head_dim) {
|
||||
case 64: launch_prefill<64>(p); break;
|
||||
case 128: launch_prefill<128>(p); break;
|
||||
default: printf("bench: unsupported D=%d\n", p.head_dim);
|
||||
}
|
||||
}
|
||||
#include "../kernels/attn_dispatchers.cuh"
|
||||
|
||||
// Warmed-up, CUDA-event timed throughput sweep over the production MMA path.
|
||||
// Reports per-call latency and effective tensor-core TFLOP/s (2 matmuls:
|
||||
// QK^T and P@V, each 2*B*Hq*ql*kl*D flops; halved for causal).
|
||||
static void bench() {
|
||||
const int cfgs[][7] = {
|
||||
{1,32,4,512,512,128,0},
|
||||
@@ -80,21 +49,21 @@ static void bench() {
|
||||
p.scale=1.0f/sqrtf((float)D);
|
||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||
|
||||
for (int i=0;i<WARMUP;i++) dispatch_prefill(p);
|
||||
auto launch = [&]() { dispatch_by_head_dim(D, [&]<int H>() { dispatch_prefill<H>(p); }); };
|
||||
for (int i=0;i<WARMUP;i++) launch();
|
||||
cudaDeviceSynchronize();
|
||||
cudaError_t err=cudaGetLastError();
|
||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return;}
|
||||
|
||||
cudaEvent_t s,e; cudaEventCreate(&s); cudaEventCreate(&e);
|
||||
cudaEventRecord(s);
|
||||
for (int i=0;i<ITERS;i++) dispatch_prefill(p);
|
||||
for (int i=0;i<ITERS;i++) launch();
|
||||
cudaEventRecord(e); cudaEventSynchronize(e);
|
||||
float ms=0; cudaEventElapsedTime(&ms,s,e); ms/=ITERS;
|
||||
|
||||
double flops = 4.0*B*Hq*(double)ql*kl*D;
|
||||
if (causal) flops *= 0.5;
|
||||
double tflops = flops/(ms*1e-3)/1e12;
|
||||
// HBM traffic: Q + O (B*Hq*ql*D each) + K + V (B*Hk*kl*D each), bf16.
|
||||
double bytes = 2.0 * (2.0*nQ + 2.0*nKV);
|
||||
double gbps = bytes/(ms*1e-3)/1e9;
|
||||
|
||||
@@ -110,6 +79,68 @@ static void bench() {
|
||||
}
|
||||
}
|
||||
|
||||
static int run_test(int B, int Hq, int Hk, int ql, int kl, int D, int causal) {
|
||||
printf("=== B=%d Hq=%d Hk=%d q=%d kv=%d D=%d causal=%d ===\n",
|
||||
B,Hq,Hk,ql,kl,D,causal);
|
||||
|
||||
size_t nQ = B*Hq*ql*D, nKV = B*Hk*kl*D;
|
||||
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
||||
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
||||
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
||||
|
||||
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||
tmp=new bf16[max(nQ,nKV)];
|
||||
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
||||
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
||||
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
||||
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
|
||||
AttentionParams<bf16> p;
|
||||
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
|
||||
p.use_mask=0; p.causal_offset=causal?0:-1;
|
||||
set_default_strides(p);
|
||||
p.scale=1.0f/sqrtf((float)D);
|
||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||
|
||||
double t0=now_ms();
|
||||
dispatch_by_head_dim(D, [&]<int H>() { dispatch_prefill<H>(p); });
|
||||
cudaDeviceSynchronize();
|
||||
double kms=now_ms()-t0;
|
||||
cudaError_t err=cudaGetLastError();
|
||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
||||
|
||||
bf16* hOut=new bf16[nQ];
|
||||
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
||||
|
||||
float* ref=new float[nQ];
|
||||
cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal ? 0 : -1);
|
||||
|
||||
float max_abs_err=0, max_rel_err=0;
|
||||
for (size_t i=0;i<nQ;i++) {
|
||||
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
||||
if(err>max_abs_err) max_abs_err=err;
|
||||
float rel=err/fmaxf(fabsf(ref[i]), 1e-8f);
|
||||
if(rel>max_rel_err) max_rel_err=rel;
|
||||
}
|
||||
const float atol=0.01f, rtol=0.01f;
|
||||
bool pass=true;
|
||||
for (size_t i=0;i<nQ;i++) {
|
||||
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
||||
if (err > atol + rtol * fabsf(ref[i])) { pass=false; break; }
|
||||
}
|
||||
printf("kernel: %.3f ms max_abs_err: %.6e max_rel_err: %.6e %s\n\n",
|
||||
kms, max_abs_err, max_rel_err, pass?"PASS":"FAIL");
|
||||
|
||||
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
|
||||
delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp;
|
||||
|
||||
return pass ? 0 : 1;
|
||||
}
|
||||
|
||||
int main() {
|
||||
const int configs[][7] = {
|
||||
{1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal
|
||||
@@ -118,59 +149,19 @@ int main() {
|
||||
{1,4,2,256,256,128,1}, // causal
|
||||
};
|
||||
int n_configs = sizeof(configs) / sizeof(configs[0]);
|
||||
int fail = 0;
|
||||
|
||||
for (int ci = 0; ci < n_configs; ci++) {
|
||||
int B=configs[ci][0], Hq=configs[ci][1], Hk=configs[ci][2];
|
||||
int ql=configs[ci][3], kl=configs[ci][4], D=configs[ci][5];
|
||||
int causal=configs[ci][6];
|
||||
printf("=== B=%d Hq=%d Hk=%d q=%d kv=%d D=%d causal=%d ===\n",
|
||||
B,Hq,Hk,ql,kl,D,causal);
|
||||
fail += run_test(B, Hq, Hk, ql, kl, D, causal);
|
||||
if (fail) break;
|
||||
}
|
||||
|
||||
size_t nQ = B*Hq*ql*D, nKV = B*Hk*kl*D;
|
||||
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
||||
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
||||
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
||||
|
||||
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||
tmp=new bf16[max(nQ,nKV)];
|
||||
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
||||
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
||||
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
||||
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
|
||||
AttentionParams<bf16> p;
|
||||
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
|
||||
p.use_mask=0; p.causal_offset=causal?0:-1;
|
||||
set_default_strides(p);
|
||||
p.scale=1.0f/sqrtf((float)D);
|
||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||
|
||||
double t0=now_ms();
|
||||
dispatch_prefill(p);
|
||||
cudaDeviceSynchronize();
|
||||
double kms=now_ms()-t0;
|
||||
cudaError_t err=cudaGetLastError();
|
||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
||||
|
||||
bf16* hOut=new bf16[nQ];
|
||||
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
||||
|
||||
float* ref=new float[nQ];
|
||||
cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal ? 0 : -1);
|
||||
|
||||
float max_err=0;
|
||||
for (size_t i=0;i<nQ;i++) {
|
||||
float d=fabsf(bf2f(hOut[i])-ref[i]);
|
||||
if(d>max_err) max_err=d;
|
||||
}
|
||||
printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err);
|
||||
|
||||
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
|
||||
delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp;
|
||||
if (fail) {
|
||||
printf("FAILED\n");
|
||||
return fail;
|
||||
}
|
||||
printf("All tests passed!\n");
|
||||
bench();
|
||||
|
||||
@@ -18,16 +18,6 @@ inline double now_ms() {
|
||||
return duration_cast<milliseconds>(steady_clock::now().time_since_epoch()).count();
|
||||
}
|
||||
|
||||
inline int compute_num_splits(int base_blocks, int tiles_total) {
|
||||
int sm_count = 0;
|
||||
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
|
||||
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
|
||||
if (n > tiles_total) n = tiles_total;
|
||||
if (n > 32) n = 32;
|
||||
if (n < 1) n = 1;
|
||||
return n;
|
||||
}
|
||||
|
||||
#define CUDA_CHECK(call) \
|
||||
do { \
|
||||
cudaError_t _e = (call); \
|
||||
|
||||
@@ -50,3 +50,4 @@ quote-style = "double"
|
||||
indent-style = "space"
|
||||
skip-magic-trailing-comma = false
|
||||
line-ending = "auto"
|
||||
exclude = ["*.md", "*.json", "*.yml", "*.yaml"]
|
||||
@@ -143,7 +143,7 @@ def print_layer_grid(results: dict[str, dict]):
|
||||
widths = [6] + [10] * len(comps)
|
||||
metric = "er_99_norm"
|
||||
|
||||
print(f"\n--- Per-Layer Effective Rank (99% energy) ---")
|
||||
print("\n--- Per-Layer Effective Rank (99% energy) ---")
|
||||
print(format_header(["Layer"] + comps, widths))
|
||||
print("-" * sum(widths))
|
||||
|
||||
@@ -173,7 +173,7 @@ def print_layer_grid(results: dict[str, dict]):
|
||||
def print_weight_stats(results: dict[str, dict]):
|
||||
groups = group_by_component(results)
|
||||
widths = [20, 12, 12, 12, 12]
|
||||
print(f"\n--- Weight Value Statistics ---")
|
||||
print("\n--- Weight Value Statistics ---")
|
||||
print(format_header(["Component", "Mean", "Std", "Min", "Max"], widths))
|
||||
print("-" * sum(widths))
|
||||
|
||||
@@ -265,7 +265,7 @@ def main():
|
||||
)
|
||||
print(f"{'=' * 70}")
|
||||
|
||||
print(f"Loading weights...")
|
||||
print("Loading weights...")
|
||||
sd = safetensors.torch.load_file(str(weights_path))
|
||||
print(f" {len(sd)} keys loaded")
|
||||
|
||||
|
||||
@@ -185,7 +185,7 @@ def choice_logprob(
|
||||
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_len
|
||||
max_len = model.config.max_position_embeddings
|
||||
if len(input_ids) > max_len:
|
||||
overflow = len(input_ids) - max_len
|
||||
input_ids = input_ids[overflow:]
|
||||
@@ -215,7 +215,6 @@ def _permute_choices(item: dict, rng: random.Random) -> tuple[dict, str]:
|
||||
positional bias (e.g. always picking B).
|
||||
"""
|
||||
letters = ("A", "B", "C", "D")
|
||||
contents = [item[k] for k in letters]
|
||||
perm = list(letters)
|
||||
rng.shuffle(perm)
|
||||
permuted = {"question": item["question"]}
|
||||
|
||||
@@ -148,7 +148,7 @@ class LossAccumulator:
|
||||
self.total += sum(losses)
|
||||
self.count += len(losses)
|
||||
if self.stream:
|
||||
clamped = [min(max(l, 0.0), self._HIST_MAX) for l in losses]
|
||||
clamped = [min(max(v, 0.0), self._HIST_MAX) for v in losses]
|
||||
idx = torch.tensor(clamped) / self._HIST_MAX * (self._HIST_BINS - 1)
|
||||
self.hist += torch.bincount(
|
||||
idx.long().clamp(0, self._HIST_BINS - 1),
|
||||
@@ -315,7 +315,7 @@ def print_stats(label: str, stats: Dict):
|
||||
)
|
||||
by_type = stats.get("by_token_type", {})
|
||||
if by_type:
|
||||
print(f"\n by token type:")
|
||||
print("\n by token type:")
|
||||
print(f" {'type':<12} {'count':>8} {'mean_loss':>10} {'ppl':>8}")
|
||||
print(f" {'-' * 12} {'-' * 8} {'-' * 10} {'-' * 8}")
|
||||
for ttype, s in by_type.items():
|
||||
|
||||
@@ -15,7 +15,7 @@ Usage::
|
||||
import argparse
|
||||
import json
|
||||
from collections import Counter
|
||||
from typing import Dict, List, Tuple
|
||||
from typing import Dict, List
|
||||
|
||||
|
||||
def _tokenize(text: str) -> List[str]:
|
||||
|
||||
+12
-12
@@ -119,15 +119,15 @@ class GenerationBenchmark:
|
||||
dtype=torch.long,
|
||||
)
|
||||
|
||||
head_dim = self.config.dim // self.config.n_heads
|
||||
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.n_layers,
|
||||
self.config.num_hidden_layers,
|
||||
batch_size,
|
||||
max_seq,
|
||||
self.config.n_kv_heads,
|
||||
self.config.num_key_value_heads,
|
||||
head_dim,
|
||||
self.device,
|
||||
self.dtype,
|
||||
@@ -136,10 +136,10 @@ class GenerationBenchmark:
|
||||
page_size = 128
|
||||
n_pages = (max_seq + page_size - 1) // page_size * batch_size
|
||||
cache = PageCache(
|
||||
self.config.n_layers,
|
||||
self.config.num_hidden_layers,
|
||||
n_pages,
|
||||
page_size,
|
||||
self.config.n_kv_heads,
|
||||
self.config.num_key_value_heads,
|
||||
head_dim,
|
||||
self.device,
|
||||
self.dtype,
|
||||
@@ -262,13 +262,13 @@ if __name__ == "__main__":
|
||||
|
||||
config = AutoRegressiveLMConfig(
|
||||
vocab_size=10000,
|
||||
dim=1536,
|
||||
n_heads=24,
|
||||
n_kv_heads=4,
|
||||
dim_ffn=6912,
|
||||
max_len=2048,
|
||||
n_layers=24,
|
||||
norm_eps=1e-5,
|
||||
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,
|
||||
)
|
||||
|
||||
benchmark = GenerationBenchmark(
|
||||
|
||||
@@ -56,7 +56,7 @@ def processor(
|
||||
print(f" {len(prompts)} prompts loaded\n")
|
||||
|
||||
if max_tokens is None:
|
||||
max_tokens = model.config.max_len
|
||||
max_tokens = model.config.max_position_embeddings
|
||||
|
||||
chunk_size = max(1, batch_size)
|
||||
|
||||
@@ -185,7 +185,10 @@ if __name__ == "__main__":
|
||||
"--max_tokens",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Maximum tokens to generate (default: model config max_len).",
|
||||
help=(
|
||||
"Maximum tokens to generate "
|
||||
"(default: model config max_position_embeddings)."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--cache_len",
|
||||
|
||||
@@ -22,9 +22,19 @@ def main():
|
||||
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
|
||||
|
||||
Pipeline(
|
||||
config=config,
|
||||
|
||||
+72
-12
@@ -1,7 +1,7 @@
|
||||
import argparse
|
||||
import os
|
||||
from functools import partial
|
||||
from typing import Any, Dict
|
||||
from typing import Any, Callable, Dict, Optional
|
||||
|
||||
import torch
|
||||
import torch.optim as optim
|
||||
@@ -12,6 +12,7 @@ from astrai.dataset import DatasetFactory, dpo_collate_fn, grpo_collate_fn
|
||||
from astrai.model import AutoRegressiveLM
|
||||
from astrai.model.components.decoder_block import DecoderBlock
|
||||
from astrai.trainer import SchedulerFactory, Trainer
|
||||
from astrai.trainer.rollout import BaseRewardModel
|
||||
|
||||
|
||||
class MuonMix(optim.Optimizer):
|
||||
@@ -101,7 +102,7 @@ def parse_args() -> argparse.Namespace:
|
||||
"--train_type",
|
||||
type=str,
|
||||
required=True,
|
||||
choices=["seq", "sft", "dpo", "grpo"],
|
||||
choices=["seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"],
|
||||
help="Train type.",
|
||||
)
|
||||
parser.add_argument(
|
||||
@@ -148,7 +149,7 @@ def parse_args() -> argparse.Namespace:
|
||||
parser.add_argument(
|
||||
"--max_grad_norm",
|
||||
type=float,
|
||||
default=None,
|
||||
default=1.0,
|
||||
help="Max gradient norm for clipping. None disables clipping.",
|
||||
)
|
||||
parser.add_argument(
|
||||
@@ -217,6 +218,39 @@ def parse_args() -> argparse.Namespace:
|
||||
default=0.0,
|
||||
help="cross_entropy function label smoothing parameter",
|
||||
)
|
||||
|
||||
# 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.",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--gradient_checkpointing",
|
||||
action=argparse.BooleanOptionalAction,
|
||||
@@ -293,8 +327,8 @@ def parse_args() -> argparse.Namespace:
|
||||
"--parallel_mode",
|
||||
type=str,
|
||||
default="none",
|
||||
choices=["none", "ddp", "fsdp"],
|
||||
help="Parallel training strategy (none, ddp, fsdp).",
|
||||
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."
|
||||
@@ -428,10 +462,19 @@ def train(
|
||||
decay_steps: int,
|
||||
**kwargs,
|
||||
):
|
||||
assert train_type in ["seq", "sft", "dpo", "grpo"]
|
||||
assert train_type 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' or 'fsdp'")
|
||||
raise ValueError(
|
||||
"--nprocs > 1 requires --parallel_mode to be 'ddp', 'fsdp', or 'fsdp2'"
|
||||
)
|
||||
|
||||
# Load config
|
||||
config_path = os.path.join(param_path, "config.json")
|
||||
@@ -439,7 +482,7 @@ def train(
|
||||
config.neftune_alpha = neftune_alpha
|
||||
|
||||
if window_size is None:
|
||||
window_size = config.max_len
|
||||
window_size = config.max_position_embeddings
|
||||
|
||||
strategy_kwargs = {
|
||||
"beta": kwargs.pop("dpo_beta"),
|
||||
@@ -449,10 +492,19 @@ def train(
|
||||
"group_size": kwargs.pop("group_size"),
|
||||
}
|
||||
|
||||
executor_kwargs = {
|
||||
"gradient_as_bucket_view": True,
|
||||
"broadcast_buffers": False,
|
||||
}
|
||||
rollout_interval = kwargs.pop("rollout_interval", 512)
|
||||
rollout_temperature = kwargs.pop("rollout_temperature", 0.7)
|
||||
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
|
||||
|
||||
executor_kwargs = {}
|
||||
if parallel_mode == "ddp":
|
||||
executor_kwargs.update(
|
||||
gradient_as_bucket_view=True,
|
||||
broadcast_buffers=False,
|
||||
)
|
||||
|
||||
model_fn = partial(create_model, config)
|
||||
dataset = DatasetFactory.load(
|
||||
@@ -510,6 +562,8 @@ def train(
|
||||
collate_fn = dpo_collate_fn
|
||||
elif train_type == "grpo":
|
||||
collate_fn = grpo_collate_fn
|
||||
elif train_type in ("online_grpo", "online_dpo"):
|
||||
collate_fn = None
|
||||
|
||||
train_config = TrainConfig(
|
||||
model_fn=model_fn,
|
||||
@@ -544,6 +598,12 @@ def train(
|
||||
extra_kwargs=strategy_kwargs,
|
||||
neftune_alpha=neftune_alpha,
|
||||
collate_fn=collate_fn,
|
||||
rollout_interval=rollout_interval,
|
||||
rollout_temperature=rollout_temperature,
|
||||
rollout_top_k=rollout_top_k,
|
||||
rollout_top_p=rollout_top_p,
|
||||
rollout_max_tokens=rollout_max_tokens,
|
||||
reward_model_fn=reward_model_fn,
|
||||
)
|
||||
|
||||
trainer = Trainer(train_config)
|
||||
|
||||
+14
-14
@@ -107,13 +107,13 @@ def test_model():
|
||||
"""Session-scoped small AutoRegressiveLM model, created once."""
|
||||
config = AutoRegressiveLMConfig(
|
||||
vocab_size=1000,
|
||||
dim=8,
|
||||
n_heads=2,
|
||||
n_kv_heads=1,
|
||||
dim_ffn=16,
|
||||
max_len=64,
|
||||
n_layers=2,
|
||||
norm_eps=1e-5,
|
||||
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"
|
||||
model = AutoRegressiveLM(config).to(device=device)
|
||||
@@ -137,13 +137,13 @@ def base_test_env(test_model, test_tokenizer):
|
||||
json.dump(
|
||||
{
|
||||
"vocab_size": 1000,
|
||||
"dim": 8,
|
||||
"n_heads": 2,
|
||||
"n_kv_heads": 1,
|
||||
"dim_ffn": 16,
|
||||
"max_len": 64,
|
||||
"n_layers": 2,
|
||||
"norm_eps": 1e-5,
|
||||
"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,
|
||||
)
|
||||
|
||||
@@ -654,6 +654,92 @@ def test_jsonl_store_sft(base_test_env):
|
||||
assert item["loss_mask"].dtype == torch.bool
|
||||
|
||||
|
||||
def test_sft_jsonl_default_messages_config(base_test_env):
|
||||
"""SFT loads a chat-style JSONL dir with no dataset_config.json.
|
||||
|
||||
Falls back to the built-in messages config: every role except
|
||||
``assistant`` is masked, loss on assistant only.
|
||||
"""
|
||||
test_dir = base_test_env["test_dir"]
|
||||
tokenizer = base_test_env["tokenizer"]
|
||||
tokenizer.set_chat_template(
|
||||
"{% for message in messages %}{{ message['role'] }}:{{ message['content'] }}\n{% endfor %}"
|
||||
)
|
||||
tokenizer_path = _save_test_tokenizer(test_dir, tokenizer)
|
||||
|
||||
data_dir = os.path.join(test_dir, "jsonl_data")
|
||||
os.makedirs(data_dir, exist_ok=True)
|
||||
records = [
|
||||
{
|
||||
"messages": [
|
||||
{"role": "system", "content": "sys"},
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": "hello"},
|
||||
]
|
||||
},
|
||||
{
|
||||
"messages": [
|
||||
{"role": "user", "content": "bye"},
|
||||
{"role": "assistant", "content": "see you"},
|
||||
]
|
||||
},
|
||||
]
|
||||
with open(os.path.join(data_dir, "data.jsonl"), "w", encoding="utf-8") as f:
|
||||
for record in records:
|
||||
f.write(json.dumps(record, ensure_ascii=False) + "\n")
|
||||
|
||||
dataset = DatasetFactory.load(
|
||||
"sft", data_dir, window_size=8, tokenizer_path=tokenizer_path
|
||||
)
|
||||
assert "sequence" in dataset.keys
|
||||
assert "loss_mask" in dataset.keys
|
||||
assert "position_ids" in dataset.keys
|
||||
assert len(dataset) > 0
|
||||
item = dataset[0]
|
||||
assert "input_ids" in item
|
||||
assert "target_ids" in item
|
||||
assert "loss_mask" in item
|
||||
assert "position_ids" in item
|
||||
assert item["loss_mask"].dtype == torch.bool
|
||||
|
||||
|
||||
def test_sft_jsonl_explicit_config_takes_priority(base_test_env):
|
||||
"""When dataset_config.json exists, it overrides the default messages config."""
|
||||
test_dir = base_test_env["test_dir"]
|
||||
tokenizer = base_test_env["tokenizer"]
|
||||
tokenizer.set_chat_template(
|
||||
"{% for message in messages %}{{ message['role'] }}:{{ message['content'] }}\n{% endfor %}"
|
||||
)
|
||||
tokenizer_path = _save_test_tokenizer(test_dir, tokenizer)
|
||||
|
||||
data_dir = _write_jsonl_dataset(
|
||||
test_dir,
|
||||
tokenizer_path,
|
||||
[
|
||||
{
|
||||
"messages": [
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": "hello"},
|
||||
]
|
||||
}
|
||||
],
|
||||
config_overrides={
|
||||
"input": {
|
||||
"sections": [{"field": "messages", "action": "$role", "template": True}]
|
||||
},
|
||||
"mask": {"user": "mask", "assistant": "train"},
|
||||
"mask_default": "mask",
|
||||
"preprocessing": {"max_seq_len": 128},
|
||||
"output": {"position_ids_mode": "doc_reset"},
|
||||
},
|
||||
)
|
||||
dataset = DatasetFactory.load(
|
||||
"sft", data_dir, window_size=8, tokenizer_path=tokenizer_path
|
||||
)
|
||||
assert "sequence" in dataset.keys
|
||||
assert "loss_mask" in dataset.keys
|
||||
|
||||
|
||||
def test_jsonl_store_pipeline_config_roundtrip(base_test_env):
|
||||
test_dir = base_test_env["test_dir"]
|
||||
config_path = os.path.join(test_dir, "dataset_config.json")
|
||||
@@ -738,7 +824,7 @@ def test_grpo_builder_preserves_response_boundaries(base_test_env):
|
||||
from tests.data.conftest import make_grpo_no_template_config
|
||||
|
||||
tokenizer = base_test_env["tokenizer"]
|
||||
tokenizer_path = _save_test_tokenizer(base_test_env["test_dir"], tokenizer)
|
||||
_save_test_tokenizer(base_test_env["test_dir"], tokenizer)
|
||||
|
||||
builder = SectionedMaskBuilder()
|
||||
config = make_grpo_no_template_config()
|
||||
@@ -851,8 +937,9 @@ def test_grpo_collate_variable_lengths():
|
||||
assert result["masks"].shape == (2, 2, 4)
|
||||
assert result["rewards"].shape == (2, 2)
|
||||
|
||||
# Check padding: item 1 prompt is length 2, padded to 3
|
||||
assert result["prompts"][1, 2] == 0
|
||||
# Prompts are left-padded so each response follows its real prompt tokens.
|
||||
assert torch.equal(result["prompts"][1], torch.tensor([0, 10, 11]))
|
||||
assert torch.equal(result["prompt_mask"][1], torch.tensor([False, True, True]))
|
||||
|
||||
# Check response content: item 0, response 0 is [4,5] padded to 4
|
||||
assert result["responses"][0, 0, 0] == 4
|
||||
|
||||
@@ -68,6 +68,28 @@ def test_chat_mask_only_assistant(chat_tokenizer, builder):
|
||||
assert len(masked) > 0
|
||||
|
||||
|
||||
def test_chat_batch_matches_single(chat_tokenizer, builder):
|
||||
config = make_chat_config()
|
||||
items = [
|
||||
{
|
||||
"messages": [
|
||||
{"role": "user", "content": "What is 2+2?"},
|
||||
{"role": "assistant", "content": "4"},
|
||||
]
|
||||
},
|
||||
{
|
||||
"messages": [
|
||||
{"role": "system", "content": "Be concise."},
|
||||
{"role": "user", "content": "Say hello."},
|
||||
{"role": "assistant", "content": "Hello."},
|
||||
]
|
||||
},
|
||||
]
|
||||
batch = builder.build_batch(items, config, chat_tokenizer)
|
||||
single = [builder.build(item, config, chat_tokenizer) for item in items]
|
||||
assert batch == single
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"mask_rules,mask_default,expect_nonzero",
|
||||
[
|
||||
@@ -152,6 +174,17 @@ def test_instruction_basic(test_tokenizer, builder):
|
||||
assert len(result["sequence"]) == len(result["loss_mask"])
|
||||
|
||||
|
||||
def test_instruction_batch_matches_single(test_tokenizer, builder):
|
||||
config = make_instruction_config()
|
||||
items = [
|
||||
{"prompt": "Translate to French: Hello", "response": "Bonjour"},
|
||||
{"prompt": "Translate to German: Hello", "response": "Hallo"},
|
||||
]
|
||||
assert builder.build_batch(items, config, test_tokenizer) == [
|
||||
builder.build(item, config, test_tokenizer) for item in items
|
||||
]
|
||||
|
||||
|
||||
def test_instruction_prompt_masked(test_tokenizer, builder):
|
||||
config = make_instruction_config()
|
||||
item = {"prompt": "hello", "response": "world"}
|
||||
@@ -363,6 +396,25 @@ def test_grpo_basic(chat_tokenizer, builder):
|
||||
assert result["rewards"] == [1.0, 0.5, 0.8, 0.2]
|
||||
|
||||
|
||||
def test_grpo_batch_matches_single(chat_tokenizer, builder):
|
||||
config = make_grpo_config()
|
||||
items = [
|
||||
{
|
||||
"prompt": [{"role": "user", "content": "What is 2+2?"}],
|
||||
"responses": ["4", "5"],
|
||||
"rewards": [1.0, 0.0],
|
||||
},
|
||||
{
|
||||
"prompt": [{"role": "user", "content": "Say hello."}],
|
||||
"responses": ["Hello", "Hi"],
|
||||
"rewards": [1.0, 0.5],
|
||||
},
|
||||
]
|
||||
assert builder.build_batch(items, config, chat_tokenizer) == [
|
||||
builder.build(item, config, chat_tokenizer) for item in items
|
||||
]
|
||||
|
||||
|
||||
def test_grpo_response_tokens_all_trained(chat_tokenizer, builder):
|
||||
config = make_grpo_config()
|
||||
item = {
|
||||
|
||||
@@ -231,3 +231,54 @@ def test_sample_with_frequency_penalty():
|
||||
)
|
||||
assert tokens.shape == (1,)
|
||||
assert 0 <= tokens[0] < logits.size(-1)
|
||||
|
||||
|
||||
def test_sample_return_logprobs_shape():
|
||||
"""``return_logprobs=True`` returns ``[batch]`` logprobs aligned to tokens."""
|
||||
logits = torch.tensor([[1.0, 2.0, 3.0], [3.0, 2.0, 1.0]])
|
||||
out = sample(logits, temperature=1.0, return_logprobs=True)
|
||||
tokens, logprobs = out
|
||||
assert tokens.shape == (2,)
|
||||
assert logprobs.shape == (2,)
|
||||
|
||||
|
||||
def test_sample_return_logprobs_nonpositive():
|
||||
"""Probabilities never exceed 1, so logprobs are always ≤ 0."""
|
||||
torch.manual_seed(0)
|
||||
logits = torch.randn(4, 50)
|
||||
_, logprobs = sample(
|
||||
logits, temperature=0.8, top_k=20, top_p=0.9, return_logprobs=True
|
||||
)
|
||||
assert torch.all(logprobs <= 1e-5)
|
||||
|
||||
|
||||
def test_sample_return_logprobs_greedy_path():
|
||||
"""Greedy decode (temperature 0) also returns logprobs."""
|
||||
logits = torch.tensor([[1.0, 5.0, 2.0]])
|
||||
tokens, logprobs = sample(logits, temperature=0.0, return_logprobs=True)
|
||||
assert tokens[0].item() == 1
|
||||
# log p(token=1) should equal log_softmax(logits)[1]
|
||||
expected = torch.log_softmax(logits.float(), dim=-1)[0, 1]
|
||||
assert torch.allclose(logprobs[0], expected, atol=1e-5)
|
||||
|
||||
|
||||
def test_sample_return_logprobs_matches_manual_computation():
|
||||
"""Returned logprob equals log_softmax(transformed_logits)[token]."""
|
||||
torch.manual_seed(1)
|
||||
logits = torch.randn(2, 30)
|
||||
tokens, logprobs = sample(logits, temperature=0.7, top_p=0.95, return_logprobs=True)
|
||||
# Recompute with the same pipeline
|
||||
from astrai.inference.sample import (
|
||||
SamplingPipeline,
|
||||
TemperatureStrategy,
|
||||
TopPStrategy,
|
||||
)
|
||||
|
||||
pipeline = SamplingPipeline([TemperatureStrategy(0.7), TopPStrategy(0.95)])
|
||||
transformed = pipeline.apply(logits.clone())
|
||||
expected = torch.gather(
|
||||
torch.log_softmax(transformed.float(), dim=-1),
|
||||
-1,
|
||||
tokens.unsqueeze(-1),
|
||||
).squeeze(-1)
|
||||
assert torch.allclose(logprobs, expected, atol=1e-5)
|
||||
|
||||
@@ -14,11 +14,11 @@ def mock_model_and_tokenizer():
|
||||
"""Create mock model and tokenizer."""
|
||||
mock_model = MagicMock()
|
||||
mock_model.config = MagicMock()
|
||||
mock_model.config.n_kv_heads = 8
|
||||
mock_model.config.n_heads = 8
|
||||
mock_model.config.dim = 128
|
||||
mock_model.config.n_layers = 2
|
||||
mock_model.config.max_len = 100
|
||||
mock_model.config.num_key_value_heads = 8
|
||||
mock_model.config.num_attention_heads = 8
|
||||
mock_model.config.hidden_size = 128
|
||||
mock_model.config.num_hidden_layers = 2
|
||||
mock_model.config.max_position_embeddings = 100
|
||||
mock_model.parameters.return_value = iter(
|
||||
[MagicMock(dtype=torch.float32, device=torch.device("cpu"))]
|
||||
)
|
||||
@@ -191,3 +191,124 @@ def test_prefill_skips_fully_cached_tasks(mock_model_and_tokenizer):
|
||||
task_id = scheduler.add_task("short prompt", stream_callback=lambda t: None)
|
||||
scheduler.stop()
|
||||
assert task_id.startswith("task_")
|
||||
|
||||
|
||||
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,
|
||||
)
|
||||
model = AutoRegressiveLM(cfg).to(device=device).eval()
|
||||
tokenizer = _Tok()
|
||||
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"
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
prompts = [[10, 20, 30], [5, 6, 7, 8]]
|
||||
results = scheduler.run_batch(prompts, max_tokens=4, temperature=1.0)
|
||||
assert len(results) == 2
|
||||
for ids in results:
|
||||
assert isinstance(ids, list)
|
||||
assert len(ids) <= 4
|
||||
assert all(0 <= i < 200 for i in ids)
|
||||
finally:
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_run_batch_return_logprobs_aligned():
|
||||
"""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]]
|
||||
results = scheduler.run_batch(
|
||||
prompts, max_tokens=5, temperature=1.0, return_logprobs=True
|
||||
)
|
||||
assert len(results) == 1
|
||||
token_ids, logprobs = results[0]
|
||||
assert len(token_ids) == len(logprobs)
|
||||
assert all(lp <= 1e-5 for lp in logprobs) # logprobs ≤ 0
|
||||
finally:
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_run_batch_respects_max_tokens():
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
prompts = [[10, 20, 30]]
|
||||
results = scheduler.run_batch(prompts, max_tokens=3, temperature=1.0)
|
||||
assert len(results[0]) <= 3
|
||||
finally:
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_run_batch_stop_id_terminates():
|
||||
"""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]]
|
||||
results = scheduler.run_batch(prompts, max_tokens=32, temperature=1.0)
|
||||
# If stop token 2 was produced, it is the last token
|
||||
if results[0] and results[0][-1] == 2:
|
||||
# No tokens after stop should exist (since we terminate)
|
||||
assert 2 not in results[0][:-1]
|
||||
finally:
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_run_batch_empty_prompts():
|
||||
"""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) == []
|
||||
finally:
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_run_batch_too_long_prompt_skipped():
|
||||
"""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
|
||||
results = scheduler.run_batch([long, [10, 20]], max_tokens=2)
|
||||
assert results[0] == []
|
||||
assert len(results[1]) <= 2
|
||||
finally:
|
||||
scheduler.stop()
|
||||
|
||||
@@ -12,13 +12,13 @@ from astrai.model.encoder import EmbeddingEncoder
|
||||
|
||||
TINY_CONFIG = dict(
|
||||
vocab_size=128,
|
||||
dim=8,
|
||||
n_heads=2,
|
||||
n_kv_heads=1,
|
||||
dim_ffn=16,
|
||||
max_len=64,
|
||||
n_layers=2,
|
||||
norm_eps=1e-5,
|
||||
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"
|
||||
@@ -42,7 +42,7 @@ def test_encoder_forward_pooling(pooling_type):
|
||||
with torch.no_grad():
|
||||
output = model(input_ids)
|
||||
|
||||
assert output.shape == (batch_size, TINY_CONFIG["dim"])
|
||||
assert output.shape == (batch_size, TINY_CONFIG["hidden_size"])
|
||||
assert not torch.isnan(output).any()
|
||||
|
||||
|
||||
@@ -60,7 +60,7 @@ def test_encoder_forward_with_padding():
|
||||
with torch.no_grad():
|
||||
output = model(input_ids, input_mask=input_mask)
|
||||
|
||||
assert output.shape == (batch_size, TINY_CONFIG["dim"])
|
||||
assert output.shape == (batch_size, TINY_CONFIG["hidden_size"])
|
||||
assert not torch.isnan(output).any()
|
||||
|
||||
|
||||
@@ -90,7 +90,7 @@ def test_encoder_from_transformer_checkpoint():
|
||||
model = _make_model()
|
||||
state_dict = model.state_dict()
|
||||
state_dict["lm_head.weight"] = torch.randn(
|
||||
TINY_CONFIG["vocab_size"], TINY_CONFIG["dim"], device=_device
|
||||
TINY_CONFIG["vocab_size"], TINY_CONFIG["hidden_size"], device=_device
|
||||
)
|
||||
|
||||
new_model = _make_model()
|
||||
|
||||
@@ -6,13 +6,13 @@ from astrai.model.transformer import AutoRegressiveLM
|
||||
|
||||
TINY_CONFIG = dict(
|
||||
vocab_size=128,
|
||||
dim=8,
|
||||
n_heads=2,
|
||||
n_kv_heads=1,
|
||||
dim_ffn=16,
|
||||
max_len=64,
|
||||
n_layers=2,
|
||||
norm_eps=1e-5,
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
@@ -58,8 +58,13 @@ CONFIGS = [
|
||||
id="gqa_qk_norm",
|
||||
),
|
||||
pytest.param(
|
||||
{**TINY_CONFIG, "attn_type": "gqa", "ffn_type": "mlp", "tie_weight": True},
|
||||
id="gqa_tie_weight",
|
||||
{
|
||||
**TINY_CONFIG,
|
||||
"attn_type": "gqa",
|
||||
"ffn_type": "mlp",
|
||||
"tie_word_embeddings": True,
|
||||
},
|
||||
id="gqa_tie_word_embeddings",
|
||||
),
|
||||
]
|
||||
|
||||
@@ -82,7 +87,11 @@ def test_model_forward(config_kwargs):
|
||||
assert "logits" in output
|
||||
assert "hidden_states" in output
|
||||
assert output["logits"].shape == (batch_size, seq_len, config.vocab_size)
|
||||
assert output["hidden_states"].shape == (batch_size, seq_len, config.dim)
|
||||
assert output["hidden_states"].shape == (
|
||||
batch_size,
|
||||
seq_len,
|
||||
config.hidden_size,
|
||||
)
|
||||
assert not torch.isnan(output["logits"]).any()
|
||||
assert not torch.isnan(output["hidden_states"]).any()
|
||||
|
||||
|
||||
@@ -19,13 +19,13 @@ from astrai.model.components.lora import (
|
||||
|
||||
MODEL_KWARGS = dict(
|
||||
vocab_size=1000,
|
||||
dim=64,
|
||||
n_heads=4,
|
||||
n_kv_heads=2,
|
||||
dim_ffn=128,
|
||||
n_layers=2,
|
||||
max_len=32,
|
||||
norm_eps=1e-5,
|
||||
hidden_size=64,
|
||||
num_attention_heads=4,
|
||||
num_key_value_heads=2,
|
||||
intermediate_size=128,
|
||||
num_hidden_layers=2,
|
||||
max_position_embeddings=32,
|
||||
rms_norm_eps=1e-5,
|
||||
)
|
||||
|
||||
|
||||
@@ -192,7 +192,7 @@ def test_inject_lora_on_moe_model():
|
||||
n_routed_experts=4,
|
||||
n_shared_experts=1,
|
||||
n_activated_experts=2,
|
||||
dim_ffn=32,
|
||||
intermediate_size=32,
|
||||
)
|
||||
inject_lora(model, r=4, alpha=8, target_modules={"up", "gate", "down"})
|
||||
assert _get_lora_count(model) > 0
|
||||
|
||||
@@ -17,13 +17,13 @@ def transformer_test_env():
|
||||
|
||||
config = {
|
||||
"vocab_size": 1000,
|
||||
"dim": 8,
|
||||
"n_heads": 2,
|
||||
"n_kv_heads": 1,
|
||||
"dim_ffn": 16,
|
||||
"max_len": 64,
|
||||
"n_layers": 2,
|
||||
"norm_eps": 1e-5,
|
||||
"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:
|
||||
@@ -45,7 +45,7 @@ def test_tie_weight_init(transformer_test_env):
|
||||
config_data = transformer_test_env["config"].copy()
|
||||
|
||||
# case 1: tie weight
|
||||
config_data["tie_weight"] = True
|
||||
config_data["tie_word_embeddings"] = True
|
||||
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_data, f)
|
||||
@@ -63,7 +63,7 @@ def test_tie_weight_init(transformer_test_env):
|
||||
assert not torch.equal(model.lm_head.weight, original_weight)
|
||||
|
||||
# case 2: not tie weight
|
||||
config_data["tie_weight"] = False
|
||||
config_data["tie_word_embeddings"] = False
|
||||
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_data, f)
|
||||
@@ -88,7 +88,7 @@ def test_model_save_load_with_tie_weight(transformer_test_env):
|
||||
config_data = transformer_test_env["config"].copy()
|
||||
|
||||
# case 1: tie weight
|
||||
config_data["tie_weight"] = True
|
||||
config_data["tie_word_embeddings"] = True
|
||||
config_path = os.path.join(test_dir, "config.json")
|
||||
|
||||
with open(config_path, "w") as f:
|
||||
@@ -108,7 +108,7 @@ def test_model_save_load_with_tie_weight(transformer_test_env):
|
||||
assert "lm_head.weight" not in model.state_dict()
|
||||
|
||||
# case 2: not tie weight (form tie-weight state dict load)
|
||||
config_data["tie_weight"] = False
|
||||
config_data["tie_word_embeddings"] = False
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_data, f)
|
||||
|
||||
|
||||
@@ -13,16 +13,16 @@ class _FakeExecutor:
|
||||
return model.state_dict()
|
||||
|
||||
|
||||
def _make_config(vocab_size=200, max_len=64):
|
||||
def _make_config(vocab_size=200, max_position_embeddings=64):
|
||||
return AutoRegressiveLMConfig(
|
||||
vocab_size=vocab_size,
|
||||
dim=16,
|
||||
n_heads=2,
|
||||
n_kv_heads=1,
|
||||
dim_ffn=32,
|
||||
max_len=max_len,
|
||||
n_layers=2,
|
||||
norm_eps=1e-5,
|
||||
hidden_size=16,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=1,
|
||||
intermediate_size=32,
|
||||
max_position_embeddings=max_position_embeddings,
|
||||
num_hidden_layers=2,
|
||||
rms_norm_eps=1e-5,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,138 @@
|
||||
"""End-to-end integration test for online DPO rollout."""
|
||||
|
||||
import os
|
||||
from functools import partial
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from astrai.config import TrainConfig
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
from astrai.trainer.rollout import BaseRewardModel
|
||||
from astrai.trainer.schedule import SchedulerFactory
|
||||
from astrai.trainer.trainer import Trainer
|
||||
|
||||
_CHAT_TEMPLATE = (
|
||||
"{% for message in messages %}"
|
||||
"{% if message['role'] == 'system' %}"
|
||||
"SYSTEM: {{ message['content'] }}\n"
|
||||
"{% elif message['role'] == 'user' %}"
|
||||
"USER: {{ message['content'] }}\n"
|
||||
"{% elif message['role'] == 'assistant' %}"
|
||||
"ASSISTANT: {{ message['content'] }}\n"
|
||||
"{% endif %}"
|
||||
"{% endfor %}"
|
||||
"{% if add_generation_prompt %}ASSISTANT: {% endif %}"
|
||||
)
|
||||
|
||||
|
||||
class InstructionDataset(Dataset):
|
||||
"""Toy instruction/input dataset for online RL rollout.
|
||||
|
||||
Each sample has an ``instruction`` and an optional ``input``; the
|
||||
RolloutGenerator renders both through the tokenizer's chat template
|
||||
so the prompt matches the SFT-trained format.
|
||||
"""
|
||||
|
||||
_SAMPLES = [
|
||||
{"instruction": "Hello", "input": ""},
|
||||
{"instruction": "Tell me a story", "input": "about dragons"},
|
||||
{"instruction": "Summarize", "input": "the article"},
|
||||
{"instruction": "Translate", "input": "to French: hi"},
|
||||
]
|
||||
|
||||
def __len__(self):
|
||||
return len(self._SAMPLES)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
return dict(self._SAMPLES[idx])
|
||||
|
||||
|
||||
class LengthRewardModel(BaseRewardModel):
|
||||
"""Rewards each response by its (non-pad) token count.
|
||||
|
||||
Enough for DPO to distinguish chosen/rejected from the rollout group.
|
||||
"""
|
||||
|
||||
def score(self, prompts, responses):
|
||||
B = len(prompts)
|
||||
G = len(responses[0]) if B else 0
|
||||
rewards = torch.zeros(B, G)
|
||||
for i in range(B):
|
||||
for g in range(G):
|
||||
rewards[i, g] = float(len(responses[i][g]))
|
||||
return rewards
|
||||
|
||||
|
||||
def instruction_collate_fn(batch):
|
||||
"""Stack a list of instruction/input dicts into a batch dict of lists."""
|
||||
return {
|
||||
"instruction": [b["instruction"] for b in batch],
|
||||
"input": [b.get("input", "") for b in batch],
|
||||
}
|
||||
|
||||
|
||||
def _model_fn(model_config):
|
||||
return AutoRegressiveLM(model_config).to(dtype=torch.float32)
|
||||
|
||||
|
||||
def _optimizer_fn(m):
|
||||
return torch.optim.AdamW(m.parameters(), lr=1e-4)
|
||||
|
||||
|
||||
def _scheduler_fn(optim):
|
||||
return SchedulerFactory.create(
|
||||
"cosine", optim, warmup_steps=1, lr_decay_steps=4, min_rate=0.05
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_online_dpo_end_to_end(base_test_env):
|
||||
"""Run one epoch of online DPO with KV-cache-backed rollout."""
|
||||
test_dir = base_test_env["test_dir"]
|
||||
device = base_test_env["device"]
|
||||
tokenizer = base_test_env["tokenizer"]
|
||||
model_config = base_test_env["transformer_config"]
|
||||
|
||||
# Equip tokenizer with a chat template so RolloutGenerator can
|
||||
# render instruction/input via apply_chat_template.
|
||||
tokenizer.set_chat_template(_CHAT_TEMPLATE)
|
||||
tokenizer.save_pretrained(test_dir)
|
||||
|
||||
model_fn = partial(_model_fn, model_config)
|
||||
optimizer_fn = _optimizer_fn
|
||||
scheduler_fn = _scheduler_fn
|
||||
|
||||
dataset = InstructionDataset()
|
||||
|
||||
train_config = TrainConfig(
|
||||
strategy="online_dpo",
|
||||
model_fn=model_fn,
|
||||
dataset=dataset,
|
||||
optimizer_fn=optimizer_fn,
|
||||
scheduler_fn=scheduler_fn,
|
||||
ckpt_dir=os.path.join(test_dir, "ckpt"),
|
||||
log_dir=os.path.join(test_dir, "logs"),
|
||||
n_epoch=1,
|
||||
batch_per_device=2,
|
||||
ckpt_interval=100,
|
||||
grad_accum_steps=1,
|
||||
random_seed=42,
|
||||
device_type=device,
|
||||
nprocs=1,
|
||||
parallel_mode="none",
|
||||
extra_kwargs={"beta": 0.1, "group_size": 2},
|
||||
rollout_interval=1,
|
||||
rollout_temperature=1.0,
|
||||
rollout_top_k=0,
|
||||
rollout_top_p=1.0,
|
||||
rollout_max_tokens=4,
|
||||
reward_model_fn=LengthRewardModel,
|
||||
collate_fn=instruction_collate_fn,
|
||||
)
|
||||
|
||||
trainer = Trainer(train_config)
|
||||
trainer.train(param_path=test_dir)
|
||||
|
||||
assert os.path.isdir(os.path.join(test_dir, "ckpt"))
|
||||
@@ -0,0 +1,355 @@
|
||||
"""Unit tests for online rollout integration in :class:`BaseStrategy`.
|
||||
|
||||
Covers the shared rollout-trigger logic in ``BaseStrategy.__call__``
|
||||
(runner injection, cache-driven refresh hook, ``step()`` callback) and
|
||||
the per-strategy ``prepare_from_rollout`` mappings for both
|
||||
:class:`GRPOStrategy` and :class:`DPOStrategy`.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
from astrai.trainer.rollout import RolloutResult
|
||||
from astrai.trainer.strategy import (
|
||||
DPOStrategy,
|
||||
GRPOStrategy,
|
||||
StrategyFactory,
|
||||
)
|
||||
|
||||
|
||||
class _FakeExecutor:
|
||||
"""Executor stub tracking ``sync_gradients`` and providing unwrap_model."""
|
||||
|
||||
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 _make_config(vocab_size=200, max_position_embeddings=64):
|
||||
return AutoRegressiveLMConfig(
|
||||
vocab_size=vocab_size,
|
||||
hidden_size=16,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=1,
|
||||
intermediate_size=32,
|
||||
max_position_embeddings=max_position_embeddings,
|
||||
num_hidden_layers=2,
|
||||
rms_norm_eps=1e-5,
|
||||
)
|
||||
|
||||
|
||||
def _make_model(device):
|
||||
cfg = _make_config()
|
||||
return AutoRegressiveLM(cfg).to(device=device), cfg
|
||||
|
||||
|
||||
def _make_frozen(model, device):
|
||||
cfg = _make_config()
|
||||
copy = AutoRegressiveLM(cfg).to(device=device)
|
||||
copy.load_state_dict(model.state_dict())
|
||||
copy.requires_grad_(False)
|
||||
copy.eval()
|
||||
return copy
|
||||
|
||||
|
||||
def _make_rollout_result(B=2, G=4, P=6, R=8, device="cpu"):
|
||||
return RolloutResult(
|
||||
prompts=torch.randint(3, 200, (B, P), device=device),
|
||||
prompt_mask=torch.ones(B, P, dtype=torch.bool, device=device),
|
||||
responses=torch.randint(3, 200, (B, G, R), device=device),
|
||||
response_mask=torch.ones(B, G, R, dtype=torch.bool, device=device),
|
||||
rewards=torch.randn(B, G, device=device),
|
||||
logprobs_old=torch.zeros(B, G, R, device=device),
|
||||
)
|
||||
|
||||
|
||||
class _RecordingRunner:
|
||||
"""Fake RolloutRunner returning a fixed result with freshness tracking.
|
||||
|
||||
Freshness is ``True`` on the first call after construction or after
|
||||
:meth:`swap_result`; ``False`` on subsequent cached calls — mirroring
|
||||
the real ``RolloutRunner`` contract without invoking generation.
|
||||
"""
|
||||
|
||||
def __init__(self, result):
|
||||
self.result = result
|
||||
self.calls = 0
|
||||
self.step_calls = 0
|
||||
self._fresh = True
|
||||
|
||||
def __call__(self, batch):
|
||||
self.calls += 1
|
||||
fresh = self._fresh
|
||||
self._fresh = False
|
||||
return self.result, fresh
|
||||
|
||||
def step(self):
|
||||
self.step_calls += 1
|
||||
|
||||
def swap_result(self, result):
|
||||
self.result = result
|
||||
self._fresh = True
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def device():
|
||||
return "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
|
||||
def _make_grpo(device, executor=None):
|
||||
model, _ = _make_model(device)
|
||||
old_model = _make_frozen(model, device)
|
||||
ref_model = _make_frozen(model, device)
|
||||
return GRPOStrategy(
|
||||
model=model,
|
||||
device=device,
|
||||
old_model=old_model,
|
||||
ref_model=ref_model,
|
||||
clip_eps=0.2,
|
||||
kl_coef=0.01,
|
||||
group_size=4,
|
||||
model_fn=lambda c=_make_config(): AutoRegressiveLM(c).to(device=device),
|
||||
executor=executor or _FakeExecutor(),
|
||||
)
|
||||
|
||||
|
||||
def _make_dpo(device, executor=None):
|
||||
model, _ = _make_model(device)
|
||||
ref_model = _make_frozen(model, device)
|
||||
return DPOStrategy(
|
||||
model=model,
|
||||
device=device,
|
||||
ref_model=ref_model,
|
||||
beta=0.1,
|
||||
reduction="sum",
|
||||
model_fn=lambda c=_make_config(): AutoRegressiveLM(c).to(device=device),
|
||||
executor=executor or _FakeExecutor(),
|
||||
)
|
||||
|
||||
|
||||
def test_factory_registers_online_aliases():
|
||||
assert StrategyFactory.is_registered("online_grpo")
|
||||
assert StrategyFactory.is_registered("online_dpo")
|
||||
assert StrategyFactory._entries["online_grpo"] is GRPOStrategy
|
||||
assert StrategyFactory._entries["online_dpo"] is DPOStrategy
|
||||
|
||||
|
||||
def test_grpo_supports_online(device):
|
||||
assert _make_grpo(device).supports_online() is True
|
||||
|
||||
|
||||
def test_dpo_supports_online(device):
|
||||
assert _make_dpo(device).supports_online() is True
|
||||
|
||||
|
||||
def test_base_strategy_prepare_from_rollout_raises_by_default(device):
|
||||
from astrai.trainer.strategy import BaseStrategy
|
||||
|
||||
class _Offline(BaseStrategy):
|
||||
def compute_loss(self, batch):
|
||||
return torch.tensor(0.0)
|
||||
|
||||
strat = _Offline(model=torch.nn.Linear(1, 1), device="cpu")
|
||||
with pytest.raises(NotImplementedError):
|
||||
strat.prepare_from_rollout(_make_rollout_result(device="cpu"))
|
||||
|
||||
|
||||
def test_base_strategy_supports_online_default_false():
|
||||
from astrai.trainer.strategy import BaseStrategy
|
||||
|
||||
class _Offline(BaseStrategy):
|
||||
def compute_loss(self, batch):
|
||||
return torch.tensor(0.0)
|
||||
|
||||
strat = _Offline(model=torch.nn.Linear(1, 1), device="cpu")
|
||||
assert strat.supports_online() is False
|
||||
|
||||
|
||||
def test_grpo_prepare_from_rollout_mapping(device):
|
||||
strat = _make_grpo(device)
|
||||
r = _make_rollout_result(device=device)
|
||||
batch = strat.prepare_from_rollout(r)
|
||||
assert batch["prompts"] is r.prompts
|
||||
assert batch["prompt_mask"] is r.prompt_mask
|
||||
assert batch["responses"] is r.responses
|
||||
assert batch["masks"] is r.response_mask
|
||||
assert batch["rewards"] is r.rewards
|
||||
|
||||
|
||||
def test_dpo_prepare_from_rollout_picks_best_worst(device):
|
||||
strat = _make_dpo(device)
|
||||
r = _make_rollout_result(B=3, G=4, R=5, device=device)
|
||||
batch = strat.prepare_from_rollout(r)
|
||||
assert batch["chosen"].shape == (3, 5)
|
||||
assert batch["rejected"].shape == (3, 5)
|
||||
assert batch["chosen_mask"].shape == (3, 5)
|
||||
assert batch["rejected_mask"].shape == (3, 5)
|
||||
idx = torch.arange(3, device=device)
|
||||
expected_best = r.responses[idx, r.rewards.argmax(dim=-1)]
|
||||
expected_worst = r.responses[idx, r.rewards.argmin(dim=-1)]
|
||||
assert torch.equal(batch["chosen"], expected_best)
|
||||
assert torch.equal(batch["rejected"], expected_worst)
|
||||
|
||||
|
||||
def test_call_without_runner_falls_back_to_compute_loss_grpo(device):
|
||||
strat = _make_grpo(device)
|
||||
batch = {
|
||||
"prompts": torch.randint(3, 200, (2, 4), device=device),
|
||||
"responses": torch.randint(3, 200, (2, 4, 6), device=device),
|
||||
"masks": torch.ones(2, 4, 6, device=device),
|
||||
"rewards": torch.randn(2, 4, device=device),
|
||||
}
|
||||
loss = strat(batch)
|
||||
assert torch.isfinite(loss).item()
|
||||
|
||||
|
||||
def test_call_with_runner_returns_finite_loss_grpo(device):
|
||||
strat = _make_grpo(device)
|
||||
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
|
||||
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
assert torch.isfinite(loss).item()
|
||||
|
||||
|
||||
def test_call_with_runner_returns_finite_loss_dpo(device):
|
||||
strat = _make_dpo(device)
|
||||
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
|
||||
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
assert torch.isfinite(loss).item()
|
||||
|
||||
|
||||
def test_call_invokes_runner_each_time(device):
|
||||
strat = _make_grpo(device)
|
||||
runner = _RecordingRunner(_make_rollout_result(device=device))
|
||||
strat.set_rollout_runner(runner)
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
assert runner.calls == 2
|
||||
|
||||
|
||||
def test_grpo_syncs_old_model_on_first_rollout(device):
|
||||
strat = _make_grpo(device)
|
||||
runner = _RecordingRunner(_make_rollout_result(device=device))
|
||||
strat.set_rollout_runner(runner)
|
||||
with torch.no_grad():
|
||||
for p in strat.model.parameters():
|
||||
p.add_(0.1)
|
||||
old_before = {k: v.clone() for k, v in strat.old_model.state_dict().items()}
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
old_after = strat.old_model.state_dict()
|
||||
synced = any(
|
||||
not torch.allclose(old_before[k], old_after[k])
|
||||
for k in old_before
|
||||
if k in old_after
|
||||
)
|
||||
assert synced
|
||||
|
||||
|
||||
def test_grpo_no_resync_when_same_cached_result(device):
|
||||
strat = _make_grpo(device)
|
||||
runner = _RecordingRunner(_make_rollout_result(device=device))
|
||||
strat.set_rollout_runner(runner)
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
assert runner.calls == 2
|
||||
assert runner.step_calls == 2
|
||||
|
||||
|
||||
def test_grpo_resync_when_new_rollout_result(device):
|
||||
strat = _make_grpo(device)
|
||||
runner = _RecordingRunner(_make_rollout_result(device=device))
|
||||
strat.set_rollout_runner(runner)
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
runner.swap_result(_make_rollout_result(device=device))
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
assert runner.calls == 2
|
||||
assert runner.step_calls == 2
|
||||
|
||||
|
||||
def test_dpo_no_sync_hook_when_new_rollout_result(device):
|
||||
"""DPO has no old_model, so ``_on_rollout_refresh`` must be a no-op.
|
||||
|
||||
We verify by ensuring no AttributeError is raised (DPO has no
|
||||
old_model) and that step is still called.
|
||||
"""
|
||||
strat = _make_dpo(device)
|
||||
runner = _RecordingRunner(_make_rollout_result(device=device))
|
||||
strat.set_rollout_runner(runner)
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
runner.swap_result(_make_rollout_result(device=device))
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
assert runner.step_calls == 2
|
||||
|
||||
|
||||
def test_step_not_called_when_sync_gradients_false(device):
|
||||
executor = _FakeExecutor(sync_gradients=False)
|
||||
strat = _make_grpo(device, executor=executor)
|
||||
runner = _RecordingRunner(_make_rollout_result(device=device))
|
||||
strat.set_rollout_runner(runner)
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
assert runner.step_calls == 0
|
||||
|
||||
|
||||
def test_step_called_when_sync_gradients_true(device):
|
||||
executor = _FakeExecutor(sync_gradients=True)
|
||||
strat = _make_grpo(device, executor=executor)
|
||||
runner = _RecordingRunner(_make_rollout_result(device=device))
|
||||
strat.set_rollout_runner(runner)
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
assert runner.step_calls == 1
|
||||
|
||||
|
||||
def test_loss_is_differentiable_grpo(device):
|
||||
strat = _make_grpo(device)
|
||||
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
|
||||
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
loss.backward()
|
||||
has_grad = any(
|
||||
p.grad is not None and p.grad.abs().sum() > 0 for p in strat.model.parameters()
|
||||
)
|
||||
assert has_grad
|
||||
|
||||
|
||||
def test_loss_is_differentiable_dpo(device):
|
||||
strat = _make_dpo(device)
|
||||
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
|
||||
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
loss.backward()
|
||||
has_grad = any(
|
||||
p.grad is not None and p.grad.abs().sum() > 0 for p in strat.model.parameters()
|
||||
)
|
||||
assert has_grad
|
||||
|
||||
|
||||
def test_ref_and_old_model_not_updated_by_backward_grpo(device):
|
||||
strat = _make_grpo(device)
|
||||
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
|
||||
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
loss.backward()
|
||||
for p in strat.ref_model.parameters():
|
||||
assert p.grad is None
|
||||
for p in strat.old_model.parameters():
|
||||
assert p.grad is None
|
||||
|
||||
|
||||
def test_ref_model_not_updated_by_backward_dpo(device):
|
||||
strat = _make_dpo(device)
|
||||
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
|
||||
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
loss.backward()
|
||||
for p in strat.ref_model.parameters():
|
||||
assert p.grad is None
|
||||
@@ -0,0 +1,382 @@
|
||||
"""Unit tests for the online rollout module."""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.inference.core.scheduler import InferenceScheduler
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
from astrai.trainer.rollout import (
|
||||
BaseRewardModel,
|
||||
RawRollout,
|
||||
RolloutGenerator,
|
||||
RolloutResult,
|
||||
RolloutRunner,
|
||||
)
|
||||
|
||||
_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 %}"
|
||||
)
|
||||
|
||||
|
||||
class FakeTokenizer:
|
||||
"""Minimal stub tokenizer with a chat template for rollout tests."""
|
||||
|
||||
stop_ids = [2]
|
||||
|
||||
def __init__(self):
|
||||
from astrai.tokenize.chat_template import ChatTemplate
|
||||
|
||||
self._chat_template = ChatTemplate.from_string(_CHAT_TEMPLATE)
|
||||
|
||||
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).decode("utf-8", errors="ignore")
|
||||
return str(ids)
|
||||
|
||||
def apply_chat_template(
|
||||
self, messages, tokenize=True, add_generation_prompt=True, **_
|
||||
):
|
||||
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 ConstantRewardModel(BaseRewardModel):
|
||||
"""Returns a constant reward for every response."""
|
||||
|
||||
def __init__(self, value: float = 1.0):
|
||||
self.value = value
|
||||
|
||||
def score(self, prompts, responses):
|
||||
B = len(prompts)
|
||||
G = len(responses[0]) if B else 0
|
||||
return torch.full((B, G), float(self.value))
|
||||
|
||||
|
||||
class BadShapeRewardModel(BaseRewardModel):
|
||||
def score(self, prompts, responses):
|
||||
return torch.zeros(len(prompts))
|
||||
|
||||
|
||||
class NonFiniteRewardModel(BaseRewardModel):
|
||||
def score(self, prompts, responses):
|
||||
B = len(prompts)
|
||||
G = len(responses[0]) if B else 0
|
||||
return torch.full((B, G), float("nan"))
|
||||
|
||||
|
||||
def _make_config(vocab_size=200, max_position_embeddings=128):
|
||||
return AutoRegressiveLMConfig(
|
||||
vocab_size=vocab_size,
|
||||
hidden_size=16,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=1,
|
||||
intermediate_size=32,
|
||||
max_position_embeddings=max_position_embeddings,
|
||||
num_hidden_layers=2,
|
||||
rms_norm_eps=1e-5,
|
||||
)
|
||||
|
||||
|
||||
def _make_model(device):
|
||||
cfg = _make_config()
|
||||
m = AutoRegressiveLM(cfg).to(device=device)
|
||||
m.eval()
|
||||
return m, cfg
|
||||
|
||||
|
||||
def _make_scheduler(model, tokenizer, max_batch_size=8, max_len=128):
|
||||
return InferenceScheduler(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=max_batch_size,
|
||||
max_seq_len=max_len,
|
||||
max_prompt_len=max_len,
|
||||
)
|
||||
|
||||
|
||||
def _make_instruction_batch(n=2):
|
||||
"""Build a batch of instruction+input prompts as lists of strings."""
|
||||
instructions = [f"Tell me about topic {i}" for i in range(n)]
|
||||
inputs = [f"context {i}" for i in range(n)]
|
||||
return {"instruction": instructions, "input": inputs}
|
||||
|
||||
|
||||
def test_raw_rollout_fields():
|
||||
r = RawRollout(
|
||||
prompts=torch.zeros(2, 4, dtype=torch.long),
|
||||
prompt_mask=torch.ones(2, 4, dtype=torch.bool),
|
||||
responses=torch.zeros(2, 3, 5, dtype=torch.long),
|
||||
response_mask=torch.ones(2, 3, 5, dtype=torch.bool),
|
||||
logprobs_old=torch.zeros(2, 3, 5),
|
||||
)
|
||||
assert r.prompts.shape == (2, 4)
|
||||
assert r.responses.shape == (2, 3, 5)
|
||||
assert r.prompt_texts == []
|
||||
assert r.response_texts == []
|
||||
|
||||
|
||||
def test_rollout_result_inherits_raw_rollout_fields():
|
||||
r = RolloutResult(
|
||||
prompts=torch.zeros(2, 4, dtype=torch.long),
|
||||
prompt_mask=torch.ones(2, 4, dtype=torch.bool),
|
||||
responses=torch.zeros(2, 3, 5, dtype=torch.long),
|
||||
response_mask=torch.ones(2, 3, 5, dtype=torch.bool),
|
||||
logprobs_old=torch.zeros(2, 3, 5),
|
||||
rewards=torch.zeros(2, 3),
|
||||
)
|
||||
assert r.rewards.shape == (2, 3)
|
||||
assert r.prompts.shape == (2, 4)
|
||||
assert r.responses.shape == (2, 3, 5)
|
||||
assert r.prompt_mask.shape == (2, 4)
|
||||
|
||||
|
||||
def test_base_reward_model_is_abstract():
|
||||
with pytest.raises(TypeError):
|
||||
BaseRewardModel()
|
||||
|
||||
|
||||
def test_constant_reward_model_shape():
|
||||
rm = ConstantRewardModel(0.5)
|
||||
out = rm.score(["a", "b"], [["x", "y", "z"], ["p", "q", "r"]])
|
||||
assert out.shape == (2, 3)
|
||||
assert torch.all(out == 0.5)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def device():
|
||||
return "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
|
||||
def _make_generator(device, **kw):
|
||||
model, _ = _make_model(device)
|
||||
tokenizer = FakeTokenizer()
|
||||
scheduler = _make_scheduler(
|
||||
model,
|
||||
tokenizer,
|
||||
max_batch_size=kw.get("max_batch_size", 8),
|
||||
max_len=kw.get("max_position_embeddings", 128),
|
||||
)
|
||||
generator = RolloutGenerator(
|
||||
scheduler=scheduler,
|
||||
tokenizer=tokenizer,
|
||||
max_tokens=kw.get("max_tokens", 8),
|
||||
group_size=kw.get("group_size", 2),
|
||||
temperature=kw.get("temperature", 1.0),
|
||||
top_k=kw.get("top_k", 0),
|
||||
top_p=kw.get("top_p", 1.0),
|
||||
)
|
||||
return generator, model
|
||||
|
||||
|
||||
def test_rollout_generator_shapes(device):
|
||||
gen, _ = _make_generator(device, group_size=3, max_tokens=5)
|
||||
batch = _make_instruction_batch(n=2)
|
||||
r = gen.generate(batch)
|
||||
assert r.responses.shape == (2, 3, 5)
|
||||
assert r.response_mask.shape == (2, 3, 5)
|
||||
assert r.logprobs_old.shape == (2, 3, 5)
|
||||
assert r.prompt_mask.shape == r.prompts.shape
|
||||
assert len(r.prompt_texts) == 2
|
||||
assert len(r.response_texts) == 2
|
||||
assert len(r.response_texts[0]) == 3
|
||||
|
||||
|
||||
def test_rollout_generator_uses_eval_and_restores_mode(device):
|
||||
gen, model = _make_generator(device, group_size=1, max_tokens=2)
|
||||
model.train()
|
||||
seen_training = []
|
||||
original = gen.scheduler.run_batch
|
||||
|
||||
def recording_run_batch(*args, **kwargs):
|
||||
seen_training.append(model.training)
|
||||
return original(*args, **kwargs)
|
||||
|
||||
gen.scheduler.run_batch = recording_run_batch
|
||||
gen.generate(_make_instruction_batch(n=1))
|
||||
assert seen_training == [False]
|
||||
assert model.training is True
|
||||
|
||||
|
||||
def test_rollout_generator_mask_matches_responses(device):
|
||||
"""Positions beyond a response's length are pad (mask False)."""
|
||||
gen, _ = _make_generator(device, group_size=2, max_tokens=6)
|
||||
batch = _make_instruction_batch(n=2)
|
||||
r = gen.generate(batch)
|
||||
for i in range(2):
|
||||
for g in range(2):
|
||||
real = r.response_mask[i, g].sum().item()
|
||||
assert r.responses[i, g, real:].sum() == 0
|
||||
if real < r.logprobs_old.size(-1):
|
||||
assert torch.all(r.logprobs_old[i, g, real:] == 0)
|
||||
|
||||
|
||||
def test_rollout_generator_logprobs_are_nonpositive(device):
|
||||
"""Behaviour-policy logprobs of sampled tokens should be ≤ 0."""
|
||||
gen, _ = _make_generator(device, group_size=2, max_tokens=4)
|
||||
batch = _make_instruction_batch(n=1)
|
||||
r = gen.generate(batch)
|
||||
for i in range(1):
|
||||
for g in range(2):
|
||||
mask = r.response_mask[i, g]
|
||||
lp = r.logprobs_old[i, g][mask]
|
||||
assert torch.all(lp <= 1e-5)
|
||||
|
||||
|
||||
def test_rollout_generator_instruction_role_mapping(device):
|
||||
"""instruction → system, input → user, output → assistant."""
|
||||
gen, _ = _make_generator(device, group_size=1, max_tokens=4)
|
||||
batch = {
|
||||
"instruction": ["Be helpful"],
|
||||
"input": ["What is 2+2?"],
|
||||
"output": ["Four"],
|
||||
}
|
||||
r = gen.generate(batch)
|
||||
text = r.prompt_texts[0]
|
||||
assert "SYSTEM: Be helpful" in text
|
||||
assert "USER: What is 2+2?" in text
|
||||
assert "ASSISTANT: Four" in text
|
||||
|
||||
|
||||
def test_rollout_generator_messages_format(device):
|
||||
"""Rollout also accepts pre-built messages."""
|
||||
gen, _ = _make_generator(device, group_size=2, max_tokens=4)
|
||||
batch = {
|
||||
"messages": [
|
||||
[{"role": "user", "content": "Hello"}],
|
||||
[{"role": "user", "content": "Goodbye"}],
|
||||
]
|
||||
}
|
||||
r = gen.generate(batch)
|
||||
assert r.responses.shape[0] == 2
|
||||
assert len(r.prompt_texts) == 2
|
||||
assert "Hello" in r.prompt_texts[0] or "USER" in r.prompt_texts[0]
|
||||
|
||||
|
||||
def test_rollout_generator_bad_batch_raises(device):
|
||||
"""Batch without messages or instruction raises a clear error."""
|
||||
gen, _ = _make_generator(device)
|
||||
with pytest.raises(
|
||||
ValueError, match="must contain either 'messages' or 'instruction'"
|
||||
):
|
||||
gen.generate({"input_ids": torch.zeros(2, 4, dtype=torch.long)})
|
||||
|
||||
|
||||
def _make_runner(device, **kw):
|
||||
generator, model = _make_generator(
|
||||
device,
|
||||
group_size=kw.get("group_size", 2),
|
||||
max_tokens=kw.get("max_tokens", 8),
|
||||
max_batch_size=kw.get("max_batch_size", 8),
|
||||
max_len=kw.get("max_position_embeddings", 128),
|
||||
)
|
||||
rm = ConstantRewardModel(1.0)
|
||||
return (
|
||||
RolloutRunner(
|
||||
generator=generator,
|
||||
reward_model=rm,
|
||||
rollout_interval=kw.get("rollout_interval", 2),
|
||||
),
|
||||
model,
|
||||
)
|
||||
|
||||
|
||||
def test_rollout_runner_shapes(device):
|
||||
runner, _ = _make_runner(device, group_size=3, max_tokens=5)
|
||||
batch = _make_instruction_batch(n=2)
|
||||
r, is_fresh = runner(batch)
|
||||
assert is_fresh
|
||||
assert r.responses.shape == (2, 3, 5)
|
||||
assert r.response_mask.shape == (2, 3, 5)
|
||||
assert r.rewards.shape == (2, 3)
|
||||
assert r.logprobs_old.shape == (2, 3, 5)
|
||||
assert len(r.prompt_texts) == 2
|
||||
assert len(r.response_texts) == 2
|
||||
assert len(r.response_texts[0]) == 3
|
||||
|
||||
|
||||
def test_rollout_runner_cache_returns_stale_flag(device):
|
||||
runner, _ = _make_runner(device, rollout_interval=10)
|
||||
batch = _make_instruction_batch()
|
||||
r1, fresh1 = runner(batch)
|
||||
r2, fresh2 = runner(batch)
|
||||
assert r1 is r2
|
||||
assert fresh1 is True
|
||||
assert fresh2 is False
|
||||
|
||||
|
||||
def test_rollout_runner_refreshes_for_different_batch(device):
|
||||
runner, _ = _make_runner(device, rollout_interval=100)
|
||||
r1, fresh1 = runner(_make_instruction_batch(n=1))
|
||||
batch2 = {"instruction": ["Different prompt"], "input": [""]}
|
||||
r2, fresh2 = runner(batch2)
|
||||
assert fresh1 is True
|
||||
assert fresh2 is True
|
||||
assert r2 is not r1
|
||||
|
||||
|
||||
@pytest.mark.parametrize("reward_model", [BadShapeRewardModel, NonFiniteRewardModel])
|
||||
def test_rollout_runner_rejects_invalid_rewards(device, reward_model):
|
||||
generator, _ = _make_generator(device, group_size=2, max_tokens=2)
|
||||
runner = RolloutRunner(generator, reward_model(), rollout_interval=1)
|
||||
with pytest.raises(ValueError):
|
||||
runner(_make_instruction_batch(n=1))
|
||||
|
||||
|
||||
def test_rollout_runner_step_triggers_new_rollout(device):
|
||||
runner, _ = _make_runner(device, rollout_interval=2)
|
||||
batch = _make_instruction_batch()
|
||||
r1, fresh1 = runner(batch)
|
||||
assert fresh1 is True
|
||||
runner.step()
|
||||
# interval=2 means trigger when _steps_since_rollout >= 2; 1 step not enough
|
||||
r2, fresh2 = runner(batch)
|
||||
assert r2 is r1
|
||||
assert fresh2 is False
|
||||
runner.step()
|
||||
# Now _steps_since_rollout == 2 -> re-rollout
|
||||
r3, fresh3 = runner(batch)
|
||||
assert r3 is not r1
|
||||
assert fresh3 is True
|
||||
|
||||
|
||||
def test_rollout_runner_clear_cache_forces_rerun(device):
|
||||
runner, _ = _make_runner(device, rollout_interval=100)
|
||||
batch = _make_instruction_batch()
|
||||
r1, _ = runner(batch)
|
||||
runner.clear_cache()
|
||||
r2, fresh2 = runner(batch)
|
||||
assert r2 is not r1
|
||||
assert fresh2 is True
|
||||
|
||||
|
||||
def test_rollout_runner_step_resets_counter(device):
|
||||
runner, _ = _make_runner(device, rollout_interval=1)
|
||||
batch = _make_instruction_batch()
|
||||
r1, _ = runner(batch)
|
||||
runner.step()
|
||||
r2, fresh2 = runner(batch)
|
||||
assert r2 is not r1
|
||||
assert fresh2 is True
|
||||
# Counter reset after rollout; second call w/o step should be cached.
|
||||
r3, fresh3 = runner(batch)
|
||||
assert r3 is r2
|
||||
assert fresh3 is False
|
||||
@@ -0,0 +1,177 @@
|
||||
import json
|
||||
import multiprocessing as mp
|
||||
import os
|
||||
import signal
|
||||
import time
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.optim as optim
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from astrai.config import TrainConfig
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
from astrai.parallel.signal_handler import register_signal_handlers
|
||||
from astrai.trainer import Trainer
|
||||
from astrai.trainer.schedule import SchedulerFactory
|
||||
from astrai.trainer.train_context import TrainContext
|
||||
|
||||
|
||||
class _PicklableDataset(Dataset):
|
||||
def __init__(self, length=200, max_length=64, vocab_size=1000):
|
||||
self.length = length
|
||||
self.max_length = max_length
|
||||
self.vocab_size = vocab_size
|
||||
|
||||
def __len__(self):
|
||||
return self.length
|
||||
|
||||
def __getitem__(self, idx):
|
||||
return {
|
||||
"input_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
|
||||
"target_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
|
||||
}
|
||||
|
||||
|
||||
def _build_model():
|
||||
config = AutoRegressiveLMConfig(
|
||||
vocab_size=1000,
|
||||
hidden_size=8,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=1,
|
||||
intermediate_size=16,
|
||||
max_position_embeddings=64,
|
||||
num_hidden_layers=2,
|
||||
rms_norm_eps=1e-5,
|
||||
)
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
return AutoRegressiveLM(config).to(device=device)
|
||||
|
||||
|
||||
class _ReadyCallback:
|
||||
def __init__(self, ready_file):
|
||||
self._ready_file = ready_file
|
||||
|
||||
def on_train_begin(self, context):
|
||||
with open(self._ready_file, "w") as f:
|
||||
f.write("ready")
|
||||
f.flush()
|
||||
os.fsync(f.fileno())
|
||||
|
||||
|
||||
def _inner_run(batch_per_device, ckpt_interval, ckpt_dir, log_dir, ready_file):
|
||||
dataset = _PicklableDataset()
|
||||
|
||||
def model_fn():
|
||||
return _build_model()
|
||||
|
||||
def optimizer_fn(m):
|
||||
return optim.AdamW(m.parameters(), lr=0.001)
|
||||
|
||||
def scheduler_fn(optim):
|
||||
return SchedulerFactory.create(
|
||||
"cosine", optim, warmup_steps=10, lr_decay_steps=10, min_rate=0.05
|
||||
)
|
||||
|
||||
train_config = TrainConfig(
|
||||
strategy="seq",
|
||||
model_fn=model_fn,
|
||||
dataset=dataset,
|
||||
optimizer_fn=optimizer_fn,
|
||||
scheduler_fn=scheduler_fn,
|
||||
ckpt_dir=ckpt_dir,
|
||||
log_dir=log_dir,
|
||||
n_epoch=1,
|
||||
batch_per_device=batch_per_device,
|
||||
ckpt_interval=ckpt_interval,
|
||||
grad_accum_steps=1,
|
||||
random_seed=42,
|
||||
device_type="cuda" if torch.cuda.is_available() else "cpu",
|
||||
)
|
||||
|
||||
trainer = Trainer(train_config)
|
||||
trainer.callbacks.insert(0, _ReadyCallback(ready_file))
|
||||
trainer.train()
|
||||
|
||||
|
||||
def _spawn_train_and_signal(ckpt_dir, sig, timeout=120):
|
||||
log_dir = os.path.join(ckpt_dir, "logs")
|
||||
ready_file = os.path.join(ckpt_dir, "ready.txt")
|
||||
|
||||
ctx = mp.get_context("spawn")
|
||||
p = ctx.Process(
|
||||
target=_inner_run,
|
||||
args=(2, 1000, ckpt_dir, log_dir, ready_file),
|
||||
)
|
||||
p.start()
|
||||
|
||||
deadline = time.time() + 30
|
||||
while time.time() < deadline:
|
||||
if os.path.exists(ready_file):
|
||||
with open(ready_file) as f:
|
||||
if f.read().strip() == "ready":
|
||||
break
|
||||
if not p.is_alive():
|
||||
break
|
||||
time.sleep(0.5)
|
||||
|
||||
assert p.is_alive(), "Training process died before becoming ready"
|
||||
|
||||
os.kill(p.pid, sig)
|
||||
p.join(timeout=timeout)
|
||||
|
||||
if p.is_alive():
|
||||
p.kill()
|
||||
p.join(timeout=5)
|
||||
|
||||
return p.exitcode
|
||||
|
||||
|
||||
def test_context_stop_flag():
|
||||
ctx = TrainContext()
|
||||
assert not ctx.stop_requested
|
||||
ctx.request_stop()
|
||||
assert ctx.stop_requested
|
||||
|
||||
|
||||
def test_register_signal_handlers():
|
||||
ctx = TrainContext()
|
||||
register_signal_handlers(ctx)
|
||||
assert not ctx.stop_requested
|
||||
os.kill(os.getpid(), signal.SIGTERM)
|
||||
assert ctx.stop_requested
|
||||
|
||||
|
||||
def test_sigterm_triggers_checkpoint_save(base_test_env):
|
||||
exitcode = _spawn_train_and_signal(base_test_env["test_dir"], signal.SIGTERM)
|
||||
assert exitcode == 0, f"Training process exited with code {exitcode} (expected 0)"
|
||||
|
||||
ckpt_dir = base_test_env["test_dir"]
|
||||
meta_files = []
|
||||
for root, dirs, files in os.walk(ckpt_dir):
|
||||
for f in files:
|
||||
if f == "meta.json":
|
||||
meta_files.append(os.path.join(root, f))
|
||||
|
||||
assert len(meta_files) > 0, f"No checkpoint meta.json found in {ckpt_dir}"
|
||||
|
||||
with open(meta_files[-1]) as f:
|
||||
meta = json.load(f)
|
||||
assert "consumed_samples" in meta
|
||||
assert meta["consumed_samples"] >= 0
|
||||
|
||||
|
||||
@pytest.mark.slow
|
||||
def test_sigint_triggers_checkpoint_save(base_test_env):
|
||||
exitcode = _spawn_train_and_signal(base_test_env["test_dir"], signal.SIGINT)
|
||||
assert exitcode == 0, f"Training process exited with code {exitcode} (expected 0)"
|
||||
|
||||
ckpt_dir = base_test_env["test_dir"]
|
||||
meta_files = []
|
||||
for root, dirs, files in os.walk(ckpt_dir):
|
||||
for f in files:
|
||||
if f == "meta.json":
|
||||
meta_files.append(os.path.join(root, f))
|
||||
|
||||
assert len(meta_files) > 0, f"No checkpoint meta.json found in {ckpt_dir}"
|
||||
Reference in New Issue
Block a user