120 Commits
Author SHA1 Message Date
ViperEkura 31d33ccdf0 chore: bump to 1.3.10 2026-07-19 16:40:27 +08:00
ViperEkura 88ec786e39 fix: memmap mode=r, tool parser json.loads, greedy decode 2026-07-19 16:38:28 +08:00
ViperEkura 663ef900fc refactor: move sample-id indexing from dataset to store
- Store owns window_size/stride and __getitem__/__len__/sample_window
- Dataset classes become thin delegators binding a Store to a train-type key mapping
- Drop BaseDataset.get_index and the RecordDataset中间类 (window死代码)
- DatasetFactory forces window_size=0 for record datasets so record semantics never get window-tainted
- token_count/num_records split the legacy len() semantics (raw stream length vs record count)
- Update tests to the new .store/.token_count API and window/record mode switching
2026-07-19 16:02:50 +08:00
ViperEkura 7d478a54db docs: update HF org from ViperEk to ViperEkura
- Replace 4 HF links in README.md and README-zh-CN.md to point to ViperEkura
- Update download.py default repo to AstrAI-V1-instruct under ViperEkura
2026-07-19 14:49:58 +08:00
ViperEkura f3eaaef842 refactor: remove redundant strategy/executor code
- Drop BaseStrategy.model_fn (stored but never read)
- Drop model_fn= passed to StrategyFactory.create in train_context
- Simplify FSDPExecutor.clip_grad_norm None branch to delegate to super()
- Remove DDPExecutor._gather_state_dict override (identical to base)
2026-07-19 12:45:58 +08:00
ViperEkura d655b65027 docs: sync architecture/dataflow/training/params with code
- dataflow.md: update DatasetFactory.load signature, stream vs record access, Store._offsets
- architecture.md: add tokenizer to Pipeline, TokenizeTransform class, RecordDataset, Streamable/Recordable mixins, fix GRPOStrategy (old_model/sync_old_model)
- training.md: DPO reduction="sum", GRPO rho_t uses pi_old, gradient_clipping always registered
- params.md: --max_grad_norm default None
2026-07-19 12:33:35 +08:00
ViperEkura 31c22dc043 refactor: deduplicate preprocessing kernel and BFD packing
- Extract shared core (mask building, primary-id extraction, tensorisation, position-id generation) to astrai/preprocessing/core.py; Pipeline and TokenizeTransform both consume it, eliminating ~60% duplicated logic
- Promote BFD _plan to module-level plan_bfd(lengths, max_len) returning pure index bins; BFDPacking.apply and evaluate_ifd._pack_bins both call it, removing the second BFD implementation
- Split Pipeline._flush (49 lines) into _inject_doc_reset_position_ids + _inject_continuous_position_ids + _to_tensors; split Pipeline.run by delegating record iteration to core.iter_raw_records
- Remove dead no-op pop/塞回 in Pipeline.run (L110-111)
2026-07-19 12:27:56 +08:00
ViperEkura 17127f8b3c fix: make tokenizer picklable for spawn multiprocessing
- ChatTemplate: defer Jinja2 compilation to cached_property, exclude compiled template from __getstate__ (its dynamic root function has __module__=None and falls back to __main__, breaking pickle)
- AutoTokenizer: bypass __getattr__ for underscore-prefixed attrs to prevent infinite recursion during unpickle when __dict__ is empty
2026-07-19 11:59:55 +08:00
ViperEkura d7695b40e3 feat: make max_grad_norm optional (None disables clipping)
- TrainConfig.max_grad_norm defaults to None
- executor.clip_grad_norm returns grad norm without clipping when None
- train.py --max_grad_norm defaults to None
2026-07-19 00:08:18 +08:00
ViperEkura fc62890e70 fix: apply chat template in DPO tokenization
- dpo_tokenize now uses tokenizer.apply_chat_template to match SFT format
- Prompt rendered with add_generation_prompt=True
- Chosen/rejected appended as assistant turn
- Remove leftover dead code from _extract_text
- Update tests to mock apply_chat_template
2026-07-19 00:00:51 +08:00
ViperEkura f433672140 fix: use sum reduction for DPO sequence logprob
- DPO requires sequence-level sum of token logprobs, not per-token mean
- mean reduction made beta*ratio_diff ~0.03 (near-zero gradient)
- loss stalled at 0.6931 because logsigmoid(0.03) has vanishing grad
- sum gives beta*ratio_diff ~10 with meaningful gradients
2026-07-18 23:48:35 +08:00
ViperEkura 7e1e5b6e6a refactor: DatasetFactory.load accepts pre-built store instance
- load(store=...) binds directly, skipping format detection/processor
- load_path now optional when store is given
- Remove redundant from_store (merged into load)
- Caller can fully control Store construction + processor setup
2026-07-18 23:23:51 +08:00
ViperEkura 553a42702d refactor: replace diamond inheritance with mixin composition
- StreamStore/RecordStore → Streamable/Recordable (stateless mixins)
- Store is sole base class, no MRO ambiguity
- H5Store/MmapStore/JsonlStore mix in both traits explicitly
- segments_are_records declared per-subclass (H5/Jsonl=True, bin=False)
- Add tests for dpo_tokenize, lazy jsonl, dual-mode H5, stream-only bin
- Remove unused _to_tensor helper
2026-07-18 23:20:41 +08:00
ViperEkura b133fc9c07 refactor: split Store into StreamStore and RecordStore
- StreamStore: fetch(begin, end, key) for stream access (SEQ/SFT)
- RecordStore: mixin with fetch_record(i, key) for record access
- H5Store/MmapStore/JsonlStore now dual-inherit both (C3 MRO)
- JsonlStore supports lazy mode via processor= (no TokenizeTransform)
- RecordDataset base class holds processor, DPO/GRPO simplified
- dpo_tokenize pure function for on-the-fly JSONL tokenisation
- DatasetFactory builds processor for jsonl+record datasets
- train.py passes tokenizer_path=param_path uniformly
- progress: len(dataset) returns sample count (stream=windows, record=records)
- json no longer auto-detected as jsonl format
2026-07-18 23:04:31 +08:00
ViperEkura b33250dc28 refactor: decouple tokenizer from Store into Transform layer
- Extract tokenization/mask/position logic from JsonlStore into TokenizeTransform
- JsonlStore now pure reader: reads JSON records, delegates to transform
- Store no longer imports tokenizer or preprocessing components
- Replace per_record param with segments_are_records class attribute
- Store subclasses declare segment semantics as format-level property
2026-07-18 21:37:31 +08:00
ViperEkura a74e5b91a3 feat: add record-mode to Store for DPO/GRPO
- Store gains fetch_record/num_records alongside stream fetch/__len__
- save_bin/load_bin support per-record offsets via record_keys param
- H5Store/MmapStore/JsonlStore all support dual stream+record access
- DPODataset/GRPODataset use fetch_record, no cross-record concat
- dpo_collate_fn + collate_fn wired through TrainConfig
- fixes attention context leakage in DPO from windowed concatenation
2026-07-18 21:02:29 +08:00
ViperEkura 28886e4241 fix: make system prompt optional across scripts
- stream_chat: default empty system_prompt, single-turn mode
- generate_batch: drop hardcoded system role
- generate.py: preserve original fields in messages branch
  and use response_key for the output column name
2026-07-18 14:10:37 +08:00
ViperEkura 9d3ccfdffc fix: incremental decode to avoid U+FFFD in streaming
- StreamDecoder buffers incomplete multi-byte sequences
- Task.decode_new_token replaces per-token decode in scheduler
- flush_remaining emits final buffered text on task finish
2026-07-18 13:05:36 +08:00
ViperEkura a24a7b4da5 perf: merge decode batch for 10x throughput
- merge all active decode tasks into single forward pass (was grouped by next_pos)
- add per-task write_positions to ContiguousCacheView for correct KV writes
- override ContiguousCache.task_cached (base returned 0, caused prefill loops)
- add --cache_len/--frequency_penalty/--rep_window to generate.py
- chunked batch processing with tqdm progress

bench (1.2B model, 128 prompts, 64 tok, batch=128):
  before: 77.2s, ~111 tok/s
  after:   7.1s, ~1210 tok/s (10.9x)
2026-07-18 08:50:46 +08:00
ViperEkura f7df02f9a3 feat: add --num_samples to batch generation script 2026-07-18 01:13:38 +08:00
ViperEkura ee450686f3 fix: add option permutation to MMLU eval
- Few-shot examples now include subject preamble (consistent format)
- Add --seed flag for option permutation (default 0, -1 to disable)
- Shuffles A/B/C/D positions per-question to neutralise positional bias
2026-07-18 00:14:34 +08:00
ViperEkura 2565755e45 refactor: switch eval datasets to HuggingFace source
- Replace GitHub/berkeley direct downloads with HF datasets API
- MMLU: cais/mmlu (all config), map val->validation split, write per-subject CSV
- HumanEval: openai/openai_humaneval
- IFEval: google/IFEval
- Enables HF_ENDPOINT mirror for faster downloads in CN
2026-07-18 00:09:31 +08:00
ViperEkura d08a92c7bd feat: add frequency penalty to inference sampling pipeline
- Add FrequencyPenaltyStrategy (logit -= penalty * count)
- Per-task rep_window for penalty history lookup
- Wire through engine, task, executor, API layer
- Add --frequency_penalty and --rep_window to stream_chat.py
- 9 unit tests for frequency penalty strategy
2026-07-17 21:28:31 +08:00
ViperEkura a1ea26d367 fix: rewrite GRPO data pipeline for offline record-level access
- process_list_field returns List[List[int]] preserving per-response boundaries
- GRPODataset rewritten to record-level __getitem__ (no windowing/stride)
- grpo_collate_fn pads variable-length responses into [B, G, R] tensors
- JsonlStore detects nested List[List[int]] and stores List[Tensor] per record
- Store._normalize skips nested-list keys from cumsum bookkeeping
- Pipeline._flush handles nested lists without cross-record flattening
- Export grpo_collate_fn from astrai.dataset
- 6 new GRPO tests + 2 updated builder tests, 114 total pass
2026-07-17 14:34:41 +08:00
ViperEkura c17aa0dc54 fix: eval script bugs and add missing features
- evaluate_mmlu: fix double few-shot injection (build_prompt no longer
  adds few-shot, apply_chat handles it once)
- evaluate_humaneval: fix pass@k k-filtering to be per-problem instead
  of using first problem's n globally; reuse ProcessPoolExecutor across
  problems; fix closure UnboundLocalError in test_one; handle None in
  report when k > n
- evaluate_ifd: remove dead code (score_plain/score_messages); add
  multi-file/directory input support with --input_path/--output_dir;
  add summary.json aggregation and --max_samples; add --dtype flag
- evaluate_ppl: add --device and --dtype flags (was hardcoded to cuda)
- evaluate_ifeval: fix docstring path (scripts/tools -> scripts/eval)
- analyze_weights: add --output JSON export; fix dead code filter
  ("_norm" not in r was always True)
2026-07-17 14:02:58 +08:00
ViperEkura b12b24eadc feat: rewrite evaluate_ppl with token-level loss and multi-file support
- Support multiple input files, glob patterns, and directory input
- Add --token_level flag: per-record token_ids + log_probs JSONL output
- Add --max_samples for random subsampling per file
- LossAccumulator: streaming mode (histogram-based percentiles, low memory) vs exact mode (full token list)
- Token type analysis (ascii/cjk/non_ascii/special) when token_level=True
- Fix token_ids/log_probs alignment (shift offset)
- Cache frozenset(stop_ids) outside loop for performance
- Aggregate stats: mean/median/ppl/p50/p90/p95/p99
- Summary JSON with all datasets in one file
2026-07-17 13:14:15 +08:00
ViperEkura cd14d53707 feat: implement bfd_split packing strategy
- BFDSplitPacking splits over-length sequences into chunks before BFD
- All keys (loss_mask, position_ids, ...) split in lockstep for alignment
- No tokens lost vs bfd which truncates over-length sequences
- Tests: token preservation, chunk alignment, short unchanged, vs bfd
2026-07-17 12:38:56 +08:00
ViperEkura e220413035 feat: support raw JSON files in dataset pipeline and JsonlStore
- detect_format now recognizes .json directories as jsonl store
- JsonlStore loads .json arrays and dicts alongside .jsonl
- tokenizer_path defaults to dataset dir when omitted
- Pipeline._iter_items handles .json files (arrays/single dict)
- Tests: detect_format, seq load, self-contained dataset dir
2026-07-17 12:20:03 +08:00
ViperEkura 84ed2327f5 feat: add --resume flag to decouple weight loading from training resumption
- Add --resume bool flag to train.py CLI
- --param_path always loads weights only by default
- --resume restores epoch, consumed_samples, optimizer & scheduler
- Checkpoint.load() now preserves full meta dict
- Update test_early_stopping to use new param_path/resume API
2026-07-16 14:23:23 +08:00
ViperEkura b14f301730 fix: init last_ckpt_step and last_log_flush_step from context.optimizer_step 2026-07-15 22:15:55 +08:00
ViperEkura 0654b4b916 refactor: template combine kernel, fix mask bug, unify dispatch
- Template combine kernel, share macros, extract entry_utils helpers
- Fix mask indexing (pass stride not pre-multiplied base)
- Remove !p.use_mask — MMA handles mask
2026-07-15 21:44:17 +08:00
ViperEkura 1f0be382ad refactor: extract load_q_mma_frags template, unify comment style
- Add load_q_mma_frags<KD>() shared template in attn_mma_utils.cuh
- Replace ~15 duplicated Q-load lines in 3 MMA kernels
- Unify section header comment style to // ---- Section ----
- Remove duplicate separator line in attn_mma_utils.cuh
2026-07-15 19:07:18 +08:00
ViperEkura bb175fda91 fix: resume optimizer LR, step display, and consumed_samples alignment 2026-07-15 08:59:52 +08:00
ViperEkura 13998da15a fix: uninitialized strides in decode test and wrong stride helper in paged test
- decode test main() missing set_default_strides → illegal memory access
- paged test used set_default_strides on PagedAttentionParams → compile error
2026-07-14 23:58:30 +08:00
ViperEkura 57729fd92d refactor: stride-based attn interface with layout and causal mask
- Replace is_causal + causal_offset with unified causal_offset (-1 = off, >=0 = first Q pos)
- Causal and mask can now coexist (was mutually exclusive)
- Add stride-based addressing for Q/KV/O (layout-agnostic, zero-copy)
- Add layout param ("bhld"/"blhd") parsed in Python, passed as int to C++
- Support 2D [batch, kv_len] and 3D [batch, q_len, kv_len] mask
- Vectorize paged KV gather in Python fallback (was per-token Python loop)
- Extract shared helpers: compute_num_splits, alloc_split_partials, dispatch_head_dim
- Unify paged_decode entry via attn_pack_paged_params
- Update mma_softmax_tile for 3D mask with per-row qrow indexing
2026-07-14 21:34:42 +08:00
ViperEkura 2c7a71a9c0 refactor: separate old policy and ref model in GRPO strategy
- Split single ref_model into old_model (importance sampling ratio) and ref_model (frozen KL regularizer)
- Move ref_model/old_model creation from strategy __init__ to TrainContextBuilder, pass as explicit parameters
- Remove periodic sync_ref_model + sync_interval; add sync_old_model for external rollout loop to call
- DPOStrategy also receives ref_model from builder
- Fix std to use unbiased=False (population std per GRPO paper)
- Remove redundant tests (test_grpo_kl_zero_at_init, test_grpo_no_sync_interval_param)
- Remove --grpo_sync_interval CLI arg
2026-07-14 20:03:45 +08:00
ViperEkura 3e0007fc91 docs : fix factory lists and MaskBuilderFactory docs
- Add MaskBuilderFactory, StoreWriterFactory, PackingStrategyFactory, PositionIdStrategyFactory to architecture design patterns
- Clarify MaskBuilderFactory three registered names (single, multi, sectioned) in preprocessing docs
2026-07-13 15:18:06 +08:00
ViperEkura b092316385 feat : add distributed checkpoint via executor checkpoint_context
- Add checkpoint_context context manager to BaseExecutor with entry/exit barrier
- Add _gather_state_dict hook overridden per executor (template method)
- DDPExecutor skips unwrap on non-rank-0 to avoid redundant state_dict gather
- FSDPExecutor uses rank0_only=True to reduce memory on non-writers
- Remove redundant rank-0 guard from Checkpoint.save and manual barrier from Callback
2026-07-13 12:27:09 +08:00
ViperEkura 9bcd696580 fix: token-level ratio and prompt masking in GRPO strategy
- Mask prompt tokens to 0 so their logprobs excluded from ratio/KL
- Switch to token-level ratio + PPO clipping via reduction='none'
- Slice response token logprobs from full sequence output
- Replace k3 KL estimator with non-negative k1 estimator
- Fix epsilon from finifo.eps (~1e-38) to 1e-8
- Remove unused 'reduction' param from GRPOStrategy.__init__
- Clarify offline batch semantics in docstring
- Add 11 unit tests for masking, advantage, KL, sync, clipping
- Sync training.md and architecture.md docs
2026-07-12 21:24:09 +08:00
ViperEkura 8f89c82d55 chore: bump version to 1.3.9 2026-07-12 21:04:20 +08:00
ViperEkura 21871197d7 refactor: use actual q_len from input, remove dead num_splits init 2026-07-12 19:48:17 +08:00
ViperEkura 4c35d36146 fix: auto-assign free port in spawn_parallel_fn to avoid EADDRINUSE 2026-07-12 19:21:56 +08:00
ViperEkura 9aca62c26c perf: remove per-element sentinel checks in softmax + fix stale comments
- Replace 4*NC8 per-element -FLT_MAX comparisons with 2 row-level pn guards; masked entries naturally underflow via expf(-FLT_MAX - nm) ≈ 0; pn only guards all-masked-row edge where nm == -FLT_MAX (exp(0)=1 not 0); ~1-3% speedup on prefill (verified via standalone CUDA bench); correctness verified: prefill 4/4, decode 3/3, paged decode 13/13
- Fix stale comments: remove false 'pre-scale Q' claim, correct occupancy numbers, remove phantom sQ from smem description
2026-07-12 15:45:08 +08:00
ViperEkura b5cdea98ad refactor: remove ineffective __launch_bounds__ from prefill kernel
- Remove MIN_BLOCKS template param and __launch_bounds__ attribute
- Profiling shows smem (not registers) is the occupancy bottleneck for D>=64, making the hint a no-op
- D=64 sees 2-4% speedup, D=128 unchanged (smem-capped at 1 block/SM)
- Update comment blocks in kernel header and both dispatch sites
- Verified correctness via standalone CUDA test (max_err ~1e-4)
2026-07-12 15:00:53 +08:00
ViperEkura 69fecaf387 perf: double-buffer KV pipeline and Q direct-to-register in decode
- Double-buffered KV (STAGES=2) for D<=128: next tile cp.async overlaps current tile MMA compute, hiding global load latency
- Q loaded directly from global into mma A-operand registers, removing sQ staging and prologue syncwarp
- Predicated cp.async unifies full and partial tile paths, eliminating scalar fallback branch
- STAGES=1 fallback for D=256 (double-buffer would exceed smem budget)
- Applied to both contiguous and paged decode MMA kernels
- ~1.27x average speedup on L20 (sm_89), zero precision loss
2026-07-12 14:14:54 +08:00
ViperEkura fd6d25ad86 refactor: extract bench_kernel and dispatch_by_head_dim into test_utils 2026-07-12 00:05:04 +08:00
ViperEkura 2c3cef1c87 feat: wire up paged decode CUDA kernel to Python extension
- Add attn_paged_decode wrapper in ops.py with gather fallback
- Register kernel in loader.py and export from __init__.py
- Extract test_utils.cuh shared by all attention unit tests
- Rename attn_paged_vs_contiguous.cu to attn_paged_decode_test.cu
- Refactor decode/prefill tests to use common bf16 helpers and cpu ref
- Fix k_cache dim check in attn_paged_decode.cu
2026-07-11 18:40:49 +08:00
ViperEkura 89ece26c25 feat: paged decode attention with split-KV (scalar + MMA)
- PagedAttentionParams merged into attn_common.h
- Scalar variant: warp-per-query-head split-KV, resolves page table per-position for the K/V shared-memory tile load
- MMA variant (sm_80+): tensor-core head-packing with cp.async, single page-table lookup per tile (BC=32 fits within page_size>=32)
- Standalone test: 14 cases across head_dim 32/64/128/256, GQA, multi-batch, both paths verified against CPU reference
2026-07-11 18:14:51 +08:00
ViperEkura 2c0b5d0b5e perf: enable MMA decode path for G=1 full attention
- inline decode_use_mma() into dispatch_decode()
- drop G>1 guard, MMA works correctly for G>=1
- decode is memory-bound, tensor cores + cp.async still win at G=1
2026-07-11 11:45:13 +08:00
ViperEkura a4ae7d17fb perf: increase decode split-K parallelism for short sequences
- Remove tiles_total/8 min-work cap that limited splits for small workloads
- Simplify decode_num_splits to only use base_blocks and tiles_total
- Short sequences now generate more blocks, improving SM utilization
2026-07-11 11:26:42 +08:00
ViperEkura 8a8550184f refactor: template AttentionParams, rename .cuh to .h
- Convert AttentionParams to a template struct supporting arbitrary types
- Rename attn_common.cuh -> attn_common.h (no CUDA-specific code remains)
- Include standard headers explicitly in each .cuh instead of via attn_common.cuh
- Allow .h files in csrc/ via .gitignore
2026-07-11 11:03:14 +08:00
ViperEkura b8b439b713 perf: fuse decode combine kernel to single-pass online-rescale reduction
- Replace 3-scan loops (mstar, lstar, acc) with 1-pass online rescale
- Halves __expf calls (num_splits vs 2*num_splits) and ml_part re-reads
- Mathematically equivalent, no change to o_part traffic or output
2026-07-11 00:24:09 +08:00
ViperEkura 41cd40363a perf: post-multiply attention scale in float instead of pre-scaling Q in bf16
- Replace bf16 pre-scale Q loading with direct 32-bit aligned bf16x2 reads
- Apply scale in float32 after Q@K^T, before online softmax
- Reduces causal max error from 2^-6 to 2^-8 with zero perf cost
2026-07-11 00:13:32 +08:00
ViperEkura d923ebe38d refactor: rename gqa_* to attn_*, split-KV for all decode paths
- Rename all csrc/kernels/gqa_*.cuh/cu to attn_*, with _split_q / _split_kv
  strategy suffix and optional _mma compute suffix
- Remove non-split MMA decode kernel, keep only split-KV path
- Convert scalar decode fallback to split-KV (o_part/ml_part + combine)
- Move combine kernel to attn_decode_split_kv.cuh (shared by both paths)
- Rename GQAParams to AttentionParams
- Update all C++ #include, PYBIND11, and Python extension references
2026-07-10 23:35:14 +08:00
ViperEkura 29b0423c4e refactor: deduplicate kernel code with shared MMA and entry-point helpers 2026-07-10 20:52:27 +08:00
ViperEkura 88f8dca2c2 perf: enlarge prefill KV tile to BC=32 for D<=128
- Kernel is latency-bound (25% occupancy), not compute/bandwidth-bound
- BC=16 wasted a cp.async wait + barrier + loop overhead per tiny tile
- Double KV tile to BC=32 for D<=128; D=256 stays 16 (64KB > 48KB smem cap)
- Retune MIN_BLOCKS per head_dim (32->6, 64->4, 128->3, 256->2)
- Result: ~6-8% faster on L20, 0.93-1.20x vs torch SDPA, correctness unchanged
2026-07-10 17:43:23 +08:00
ViperEkura 9027fdc546 test: bench production MMA attention path with FLOP/s and bandwidth 2026-07-10 16:58:34 +08:00
ViperEkura cbd140340d perf: load prefill Q fragments directly from global (drop sQ staging)
- read the 8 Q elements each lane needs straight from global into the mma
  A-operand layout, pre-scaled, instead of staging through shared sQ
- removes the sQ smem area (20KB->16KB) and the serialized per-warp prologue
  with its WARPS __syncthreads barriers
- result vs torch SDPA: prefill 0.70-0.82x -> 0.85-1.00x (matches torch at
  seq=128; 2048 1.251->1.159ms), correctness unchanged across head dims
2026-07-10 12:27:45 +08:00
ViperEkura 988e01314d perf: pipeline prefill MMA kernel (double-buffered K/V + packed stores)
- double-buffer K/V one tile ahead via cp.async to overlap load with tensor-core math (ncu long_scoreboard 2.12->0.53)
- reorder wait->barrier->prefetch so one __syncthreads/tile covers both cross-warp publish and buffer-reuse (was two)
- add predicated cp_async_16_pred (src-size=0 zero-fills OOB) to unify full/partial tiles, dropping the scalar fallback
- halve BC to 16 to keep 3 blocks/SM despite the doubled smem
- pack adjacent bf16 output into one 32-bit STG, removing the uncoalesced scalar-store penalty (14%->5% sectors)
- result vs torch SDPA: prefill 0.61-0.78x -> 0.70-0.82x, spills eliminated
2026-07-10 12:19:14 +08:00
ViperEkura 7ba43a7c6f perf: add split-K (FlashDecoding) to decode MMA kernel
Decode has only batch*kv_head independent tasks, so the grid was tiny (e.g. 16 blocks) leaving most SMs idle (ncu: 0.04 waves/SM, 11% DRAM).

- Partition KV across gridDim.z blocks emitting unnormalised (O, m, l) partials, reduced by a new combine kernel
- Choose split count to fill the device (~2 blocks/SM), capped by tile count and 32; fall back to single-pass direct-write when batch*kv_head already saturates the SMs
- Refactor decode dispatch into named helpers, de-duplicate scalar fallback

Result: now DRAM-bound at 63% (99->543 GB/s), 2.1-2.5x over torch SDPA in the low-parallelism regime, on par at high parallelism
2026-07-10 11:43:18 +08:00
ViperEkura dea59f7e1d fix: restore resident Qa to fix sQ overwrite bug
Moving Qa ldmatrix into the tile loop caused warps 0-2 to read
warp 3's Q data from sQ (only the last warp's data survives the
serialized load loop). Reverted to loading Qa during the init phase
and keeping it resident; __launch_bounds__ still forces 128 regs
(33% occupancy) with spill to local memory.
2026-07-10 00:49:58 +08:00
ViperEkura 85dc771460 perf: reduce MMA kernel registers, switch to static smem
- Move Qa[KD][4] into tile loop (reload from sQ per tile)
  cutting ~32 resident registers for HEAD_DIM=128
- Replace extern __shared__ with static template-sized smem
  (no cudaFuncSetAttribute or dynamic allocation needed)
- Add __launch_bounds__ with MIN_BLOCKS param, dispatch by HEAD_DIM
  (hd=128→4, hd=64→6, hd=32→6)
- Remove dynamic smem from scalar kernel and C test
- Result: hd=128 168→128 regs, 25%→33% occupancy
2026-07-10 00:39:47 +08:00
ViperEkura 2c5629b81d docs: fix documentation errors across README and assets/docs
- Correct training CLI args: remove non-existent --adamw_beta1/2, fix --weight_decay
- Fix optimizer section: document MuonMix instead of plain AdamW
- Fix inference.md decode phase description (all groups, not largest)
- Fix dataflow.md H5Store description (no share_memory_ in code)
- Fix architecture.md class diagram: remove Task.stream_callback,
  EncoderConfig.use_gated_attention; update stop_ids docs
- Fix --min_rate default value description to match code (0.01)
- Update all document timestamps to 2026-07-09
2026-07-09 10:09:53 +08:00
ViperEkura 841a582b28 refactor: split mask builder by single/multi output
- Extract SingleOutputMaskBuilder for SFT and pretrain configs
- Extract MultiOutputMaskBuilder for DPO and GRPO configs
- Keep SectionedMaskBuilder as backward-compatible facade
- Register "single" and "multi" names in MaskBuilderFactory
- Add parity and rejection tests for concrete builders
2026-07-08 21:18:34 +08:00
ViperEkura c8567a6f65 fix: exclude embedding, lm_head, bias, and norm params from Muon optimizer, use AdamW 2026-07-08 19:42:20 +08:00
ViperEkura 8035be9b1f fix: make MuonMix inherit from torch.optim.Optimizer 2026-07-08 17:03:05 +08:00
ViperEkura e9b03f4fca perf: apply cp.async, XOR swizzle, pre-scaled Q to decode MMA kernel
Decode MMA kernel previously used scalar global→shared loads with
LD=HEAD_DIM+8 padding and per-tile scale multiply. This commit brings it
in line with the prefill MMA kernel (which already had these optimizations):

- cp.async K/V loads (bypasses registers, halves load instructions)
- XOR swizzle: LD=HEAD_DIM instead of HEAD_DIM+8 (zero waste smem)
- Pre-scale Q during load (removes per-tile scale multiply in softmax)
- Clean up prefill MMA kernel comments (no code change)

~2x speedup on decode (0.47ms→0.24ms at seq_len=512)
2026-07-08 16:15:14 +08:00
ViperEkura fd65b9bc23 feat: support HEAD_DIM=32 and split extension into loader/ops
- add case 32 to decode/prefill dispatch switch
- fix swiz_col out-of-bounds for HEAD_DIM=32: XOR mask now limited to chunk count (3 for 32, 7 for >=64) instead of always 7, which produced column offsets >= LD=32 and corrupted shared memory
- restructure decode dispatch to #ifndef/#else/#endif matching prefill
- split astrai/extension/__init__.py into loader.py (kernel .so discovery) and ops.py (wrapper functions + torch SDPA fallback); __init__.py now re-exports the public API
2026-07-08 14:14:11 +08:00
ViperEkura 9ebaea840f perf: cp.async K/V loads, shared sQ staging, causal skip, XOR swizzle
- cp.async global→shared for K/V full-tile loads, eliminates 99.6% of shared-store bank conflicts (612K→2.7K per ncu)
- add cp_async_16/commit/wait_all/wait_group<N> helpers in mma utils
- shared sQ staging (single area, serialized per-warp load), cuts smem from (2*BC + WARPS*BR)*LD to (2*BC + BR)*LD bf16
- pre-scale Q by attention scale during Q load, removes per-tile scale multiply in softmax loop
- causal tile skipping: block-level early break + warp-level skip
- scalar fallback only for last partial tile
- XOR swizzle (swiz_col) at 8-bf16 chunk granularity, eliminates ldmatrix bank conflicts without LD padding, LD=HEAD_DIM (zero smem waste), saves 1280 bytes/block vs HEAD_DIM+8 padding
2026-07-08 12:30:32 +08:00
ViperEkura 6adc221c10 refactor: extract shared MMA utils into gqa_mma_utils.cuh
- Move mma16816, ld2, pk2, pkb, ldmatrix_x4/x2/x2_trans to shared header
- gqa_prefill_attn_mma.cuh and gqa_decode_attn_mma.cuh both include it
2026-07-07 23:01:15 +08:00
ViperEkura 9e63cb9ed0 feat: MMA head-packing decode kernel with scalar fallback dispatch
- Add gqa_decode_attn_mma.cuh for tensor-core decode path
- Add dispatch_decode<> selecting MMA vs scalar based on G and mask
- Add TORCH_CHECK for unsupported head_dim instead of silent scalar launch
2026-07-07 22:56:02 +08:00
ViperEkura 4225518cf3 perf: add fast-math and vectorization nvcc/cxx build flags
- centralize CXX_FLAGS/NVCC_FLAGS in csrc/build.py as single source
- add --use_fast_math, --ptxas-options=-O3,-v, --extra-device-vectorization
- add -march=native -funroll-loops host flags
- setup.py reads shared cxx_flags/nvcc_flags from registry
- sync pure-C test build commands with new flags
2026-07-07 22:28:32 +08:00
ViperEkura c50adbaac0 feat : replace AdamW with MuonMix (Muon + AdamW) optimizer
- Muon for 2D matrix params, AdamW for 1D (norm/bias/embed)
- MuonMix wrapper handles combined step/zero_grad/state_dict
- New CLI args: weight_decay, muon_momentum, muon_nesterov, muon_ns_steps, muon_adjust_lr
- Removed adamw_beta1/adamw_beta2/adamw_weight_decay
- Moved optimizer/strategy params from signature to **kwargs
2026-07-07 14:10:36 +08:00
ViperEkura 536dbc0c9a fix: set tqdm postfix before update so first step shows metrics 2026-07-07 00:14:13 +08:00
ViperEkura 4af7acd449 fix: support single .h5 file loading in load_h5 2026-07-07 00:11:16 +08:00
ViperEkura 53ed52b4b8 refactor: extension dispatch layer with CUDA/torch fallback
- Add gqa_decode_attn/gqa_prefill_attn dispatch functions
- Internal _available/__modules with underscore prefix
- CUDA kernel path with F.scaled_dot_product_attention fallback
- GQA head expansion in fallback path
2026-07-06 21:07:16 +08:00
ViperEkura f1cc7cedce feat: ldmatrix + smem padding for mma prefill kernel
- replace scalar fragment loads with ldmatrix.sync.x4/x2
- add smem row-stride padding (LD = HEAD_DIM + 8) to eliminate 8-way bank conflicts from HEAD_DIM being a 32-bank multiple
- switch build flag from positive to negative: -DASTRAI_NO_MMA for pre-sm_80 only; mma is the default path
- vectorize scalar path smem loads with float4 ld8
- fix pure-C test configs for ld8 alignment
2026-07-06 20:55:22 +08:00
ViperEkura ddc4bd1cf6 feat: tensor-core mma prefill with build-time dispatch
- add register-resident flash-attention kernel using mma.sync.m16n8k16
- dispatch mma vs scalar at build time: pre-sm_80 defines
  -DASTRAI_NO_MMA, else defaults to mma
- scalar path vectorized with float4 smem loads (ld8)
2026-07-06 20:33:24 +08:00
ViperEkura cc36530c73 perf: group-split register-blocking gqa_prefill kernel
- one query row per group of G=8 lanes, each owning HEAD_DIM/G dims of qreg[]/acc[] in registers
- removes full 32-lane warp_reduce_sum; S dot reduces over only G lanes
- templated on <HEAD_DIM,G,ROWS,P_BC>, block=(G,ROWS)=(8,32)
- per-group shuffle mask so causal loop-bound divergence doesn't deadlock the shuffle
- update pure-C test to the templated launch
2026-07-06 18:33:08 +08:00
ViperEkura 11fa807cfc fix: correct prefill mask index, unify GQA kernel interface
- Fix mask indexing: batch*q_len*kv_len -> batch*kv_len
- Add csrc/kernels/gqa_common.cuh with shared GQAParams struct
- Unify decode/prefill Python API: both accept (q,k,v,mask=None,...)
- Decode now supports optional mask, is_causal, causal_offset, scale
- Rename struct fields: B->batch, Hq->q_head, Hk->kv_head, D->head_dim
- Use py::arg() for correct None/defaults handling in pybind11
- Update pure C tests and build instructions (-arch=sm_89)
2026-07-06 17:21:23 +08:00
ViperEkura bcdd93e0eb feat: split kernel defs from bindings, add prefill tiled kernel and pure C tests
- Split .cuh/.cu for gqa_decode_attn and gqa_prefill_attn
- gqa_prefill_attn: tiled shared-memory K/V, fused load, compute-opt, mask support
- Add pure C tests under csrc/tests/ for fast nvcc-only iteration
- Update .gitignore for build artifacts
2026-07-06 16:14:55 +08:00
ViperEkura 579b8c3129 fix: correct gqa_decode_attn reduction + add gqa_prefill_attn
- gqa_decode_attn: rewrite to per-KV-head, K in smem
- gqa_prefill_attn: new kernel for Q_len > 1 with GQA
2026-07-06 13:45:18 +08:00
ViperEkura d7da51569f docs: update install instructions in EN/CN README 2026-07-06 12:25:36 +08:00
ViperEkura e8e228d035 feat: add optional CUDA kernel system (csrc/) + fused GQA decode attention
Structure:
  csrc/               -- .cu sources + build.py registry
  astrai/extension/   -- compiled .so + __init__.py (import dispatcher)
  setup.py            -- CUDAExtension from csrc/build.py REGISTRY

Control: CSRC_KERNELS=true|false env var at install time.
Fallback: astrai.extension.available dict for runtime detection.
2026-07-06 12:09:58 +08:00
ViperEkura 2579658e15 chore : shields release badge from /release to /tag 2026-07-05 20:34:40 +08:00
ViperEkura f0cd0134c6 fix : update benchmark for v1.3.8 cache API, add argparse and cache type switch
- Replaced old KVCage API with PageCache/ContiguousCache
- Added --cache contiguous|paged switch for decoding comparison
- Added argparse for all params (batch/prompt/gen/device/dtype)
- Fixed PageCache decode crash by extending pages for full sequence
2026-07-05 20:30:26 +08:00
ViperEkura abb96996f8 docs : sync 6 doc files to actual code
- architecture.md: removed TrainConfig.log_interval, split KVCache into
  PageCache/ContiguousCache with CacheView/PageCacheView/ContiguousCacheView,
  added JsonlStore, fixed GradientCheckpointingCallback type,
  CheckpointCallback typo, ProgressBarCallback hooks
- training.md: added position_ids to SFT keys, fixed callback hook table,
  removed merged ValidationCallback
- inference.md: documented ContiguousCache default vs PageCache paged
- dataflow.md: added JsonlStore to storage backends and format detection
- params.md: removed nonexistent --log_interval
- preprocessing.md: updated timestamp
2026-07-05 19:35:18 +08:00
ViperEkura bbe6ff2d8f release : v1.3.8
- refactor: 重写 IFD 评估为三层架构,引入 BFD 装箱与自定义 attention mask 批处理打分
- refactor: 重写 HumanEval 评估为函数式流水线,修复测试超时与动态 pass@k
- perf: 替换 paged KV cache 为 ContiguousCache,解码所有 group
- feat: 新增 ROUGE 评估脚本、JSONL 数据集 store、stream_chat 参数
- fix: 修复 IFD token-set 不对称、SFT position_ids 默认值、文档边界保留
2026-07-05 19:12:33 +08:00
ViperEkura db9b39b084 fix: resolve IFD token-set asymmetry and support single-token answers
- Sentinel-anchored unconditional pass: both branches now predict the same N response tokens
- Single-token responses (rl=1) fully supported
- ctx_len tracked per sample; skip_reason replaces silent None
- --per_token flag for per-token IFD breakdown
2026-07-05 17:48:26 +08:00
ViperEkura 849e1e00a3 refactor: clean up inference design patterns
1. KVCache base: add default task_cached/task_record_hashes, remove getattr from scheduler
2. Remove page_size param from scheduler constructor (ContiguousCache-only)
3. InferenceEngine expose cache param for KVCache injection
4. Rename page_cache -> kv_cache in Executor
5. Move stream_callback from Task to TaskManager._callbacks dict
6. TaskManager.clear_queues clears callbacks
2026-07-05 11:41:54 +08:00
ViperEkura 5416c2e8fb perf: replace paged KV cache with contiguous ContiguousCache, decode all groups
- Add KVCache/CacheView abstract base classes in cache.py
- Add ContiguousCache (contiguous per-slot buffer, default) alongside PageCache (paged, renamed from old KVCache)
- Merge make_table_tensor + bind into bind_tasks on KVCache interface
- Remove task_cached/task_record_hashes from base class (PageCache-only)
- Scheduler: decode all position groups instead of just the largest (eliminates 63% group skip rate)
- Scheduler: accept optional cache param for swapping implementations
- Model layer type hints use CacheView base class
- Batch 1-32: 1-7% speedup from eliminating Storage.gather overhead
- All 183 inference tests pass
2026-07-05 11:34:36 +08:00
ViperEkura 599a51f4f7 fix: reliable test timeout, separate generate/test phases, dynamic pass@k
- Replace SIGALRM+exec() with subprocess.run(timeout=) for test execution
- Add --test_only flag to skip generation and test existing completions
- Add --generate_only flag for generation-only runs
- Derive pass@k values from num_samples (filter k > n)
- Support loading completions from array JSON (not just JSONL)
2026-07-05 08:47:30 +08:00
ViperEkura 17d6eaa2f2 refactor: rewrite humaneval evaluation with functional pipeline design
- fix KeyError race condition in inference cache touch()
- EvalConfig dataclass for centralized configuration
- load->generate->extract->test->score->report pipeline
- two-phase generation+testing for max GPU utilization
- signal-based SIGALRM timeout protection for code exec
- suppress subprocess stdout/stderr pollution
2026-07-05 07:58:28 +08:00
ViperEkura 2d908639e9 feat : add ROUGE evaluation script (manual impl, no deps)
- ROUGE-1/2 via n-gram overlap (Counter)
- ROUGE-L via LCS (DP)
- CLI: python scripts/eval/evaluate_rouge.py --data_path ... --output ...
- Library: compute_rouge(ref, cand) -> dict of precision/recall/f1
2026-07-05 01:15:01 +08:00
ViperEkura c7158418dd perf: add BFD bin-packing and custom attention mask to IFD batch scoring 2026-07-04 18:58:13 +08:00
ViperEkura 4d3c9341c1 refactor: rewrite IFD evaluation with clean three-layer architecture 2026-07-04 18:33:51 +08:00
ViperEkura 4e508afa2d fix : SFT pipeline position_ids default & doc boundary preservation
- change position_ids_mode default from "none" to "doc_reset" so SFT preprocessing always generates position_ids (was causing dataset load KeyError)
- generate per-doc position_ids before packing (doc_reset mode), preserving document boundaries for BFD packing (cross-doc attention leak fix)
- change _align_bucket padding from [1] to [0] to avoid accidentally training on loss_mask padding
2026-07-04 15:59:11 +08:00
ViperEkura 8999ca89b8 feat: add JSONL dataset store with on-the-fly tokenization
- Add JsonlStore registered under "jsonl" in astrai/dataset/storage.py
- Reuse PipelineConfig schema for JSONL dataset configuration
- Update detect_format to recognize JSONL directories and files
- Move save_h5/load_h5/save_bin/load_bin to astrai/serialization
- Split astrai/serialization.py into checkpoint/dataset submodules
- Add tests for JSONL detection, seq/SFT stores, and config roundtrip
2026-07-04 15:42:33 +08:00
ViperEkura 1adca39cd8 fix: handle long sequences and optimize IFD computation 2026-07-04 08:35:45 +08:00
ViperEkura 204873fa2f fix: handle long sequences and optimize IFD computation 2026-07-04 08:23:32 +08:00
ViperEkura a5c1de6b1b feat: add model_path temperature top_p top_k max_tokens system_prompt args to stream_chat 2026-07-04 07:33:32 +08:00
ViperEkura 27524ad085 fix: reset sampler iter at epoch end so progress bar shows total after first epoch 2026-07-04 06:35:55 +08:00
ViperEkura 27d1921d9c fix: scheduler division-by-zero, loss_mask bool
- schedule.py: guard warmup_steps/lr_decay_steps against zero
- strategy.py: use ~loss_mask instead of loss_mask==0 on bool tensor
2026-07-03 22:04:55 +08:00
ViperEkura 70c0e5de90 refactor: merge validation into MetricCallback, simplify progress bar to optimizer steps
- Remove separate ValidationCallback, merge into MetricCallback
- Progress bar now tracks optimizer steps instead of micro-steps
- Remove unused log_interval config field and CLI flag
- Fix validation all_reduce: use SUM(loss, count) instead of AVG
- Simplify metric logging: always log every optimizer step
- Add grad_norm display to progress bar
2026-07-03 21:43:08 +08:00
ViperEkura dfb151537b fix: ForwardRef._evaluate Python 3.12 compatibility 2026-07-03 18:41:19 +08:00
ViperEkura 500c605fad fix: unify scheduler min_rate default to 0.01, clamp WSD warmup 2026-07-03 17:52:23 +08:00
ViperEkura dc9faca3b1 fix: align docs with actual code (40+ inconsistencies)
- Remove nonexistent Muon class from architecture diagram
- Fix Checkpoint/TrainConfig/TrainContext field names (iteration -> consumed_samples, start_batch -> start_samples)
- Add missing fields: neftune_alpha, val_split, grad_norm, optimizer_step, tool_calls/tools
- Fix CLI param defaults: --log_interval 1, --metrics [loss,lr,grad_norm], --start_samples
- Add missing scheduler CLI params; remove nonexistent --num_workers from preprocess docs
- Fix inference SSE format, stats response keys, error codes to match actual server output
- Fix preprocessing docs: BOS once, shard_0000 layout, from_json->from_file, GRPO prompts_mask
- Fix dataflow detect_format/_normalize descriptions; correct callback order in training.md
2026-06-30 20:47:23 +08:00
ViperEkura aabb0d83e9 refactor : replace iteration with consumed_samples
- Replace context.iteration with consumed_samples (global sample count)
- Add optimizer_step property derived from consumed_samples
- Checkpoint meta.json stores consumed_samples, drops iteration
- CLI --start_batch renamed to --start_samples (per-rank samples)
- Checkpoint dir naming: epoch_X_step_Y instead of epoch_X_iter_Y
- Metric log entries use step and consumed_samples fields
- Backward compat removed (old iteration checkpoints unsupported)
2026-06-30 18:42:42 +08:00
ViperEkura 44579ea6dc refactor : metric 日志改为以 optimizer step 为单位,默认每步记录
- log_interval 默认 100 -> 1,语义从 batch iteration 改为 optimizer step
- step 指标从 on_batch_end 移到 on_optimizer_step,不受梯度累积影响
- JSONL 条目新增 step 字段,保留 iter
- flush 落盘仍在 on_batch_end
2026-06-30 15:12:31 +08:00
ViperEkura 0f1fcb079f refactor : grad_norm 指标简化,clip_grad_norm 移至 executor
- metrics 默认加入 grad_norm,移除 grad_std/max/min/mean/nan_num
- grad_norm 默认返回总 L2 范数,per_param=True 返回各参数范数
- clip_grad_norm 从 callback 移至 BaseExecutor/FSDPExecutor
- FSDPExecutor 覆盖为 model.clip_grad_norm_() 保证分布式正确
- ctx_get_grad_norm 改为读取 context.grad_norm
2026-06-30 14:59:43 +08:00
ViperEkura 84d4769163 feat: SVD 有效秩/权重统计分析脚本 2026-06-29 21:39:22 +08:00
ViperEkura bf09a35c95 feat: optimizer 参数分组,bias/norm 不做 weight decay 2026-06-27 16:30:34 +08:00
ViperEkura 6715461a36 chore : 升级 torch 2.11.0+cu128,移除自定义 Muon,修复 gloo device_id
- torch 2.7.1-cu126 升级至 2.11.0-cu128,numpy 2.3.2 升级至 2.4.4
- 移除 astrai/trainer/optim.py,改用 torch.optim.Muon
- parallel setup: gloo 后端不再传递 device_id,单卡多进程不再报错
2026-06-27 16:10:37 +08:00
ViperEkura b4587c5d08 refactor : metric_logger 改用事件类型 (type=step/validation/epoch)
- 每种事件独立 schema,不再混入 null 字段
- 回调顺序 validation 移到 metric_logger 之前,确保 on_optimizer_step 先跑
- 用内部 _last_val_loss 代替 TrainContext.last_val_iter 判断新验证
- 修复 factory.py 未使用导入、evaluate_ifeval.py 多余 f 前缀
2026-06-25 17:18:20 +08:00
ViperEkura 88ec63121d feat : GPT-2 residual scaling weight init
- Linear: normal(0, init_std) replaces kaiming_uniform_(a=sqrt(5))
- o_proj / mlp.down: init_std = 0.02 / sqrt(2 * n_layers)
- MoE: expert down scaled by 1/sqrt(1/n_shared + 1/K)
- Embedding: normal(0, 0.02), unchanged
2026-06-25 15:08:31 +08:00
ViperEkura 01d2da2893 feat : 训练支持 --schedule_type 及对应调度器参数
- --schedule_type 可选 cosine/sgdr/wsd,默认 cosine
- --min_rate 统一控制最小 LR 比率
- --cycle_length / --t_mult 用于 sgdr
- --stable_steps / --decay_steps 用于 wsd,自动计算默认值
2026-06-22 10:35:56 +08:00
ViperEkura 25d4ea3f91 refactor : 压缩测试代码,消除重复
- fixture 替代重复实例化和 tokenizer 落盘
- parametrize 合并同构测试
- helper 消除 save_h5 + DatasetFactory.load 样板
- 净减 272 行
2026-06-19 14:54:39 +08:00
ViperEkura 39985840c7 refactor : neftune_alpha 在 Embedding 构造时传入,由模型配置链路负责
- BaseModelConfig 添加 neftune_alpha 字段 (默认 0.0)
- Embedding.__init__ 接受 neftune_alpha 参数,不再外部 set
- AutoRegressiveLM / EmbeddingEncoder 从 config 传入 neftune_alpha
- train.py 将 CLI 参数注入 config 后再创建模型
- TrainContextBuilder 移除 neftune 设置(不再是其职责)
2026-06-19 14:23:27 +08:00
ViperEkura b1adc40cfb refactor : 将 config 对象直接传给 DecoderBlock,替代 16 个独立参数
- DecoderBlock.__init__ 改为 (config, layer_id),内部用 asdict
  展开字段给 AttnFactory/FFNFactory,factory 按 __init__ 签名自动过滤
- EncoderConfig 补充 attn_type 和 ffn_type 字段
- 314 个测试全部通过
2026-06-19 14:15:33 +08:00
ViperEkura 7348bac6ab fix: 规范 generate.py 命令行接口
- generate.py 清理描述文字,help 统一标注默认值
- max_tokens 默认改为 None,回退 model config max_len
- evaluate_ppl.py 同步清理描述文字
- params.md 同步 max_tokens 默认值
2026-06-19 14:03:02 +08:00
110 changed files with 9913 additions and 2621 deletions
+71
View File
@@ -0,0 +1,71 @@
name: Release
on:
push:
tags:
- "v*"
jobs:
build-pure:
name: Build pure-Python wheel
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: "3.12"
- name: Build wheel (no CUDA)
run: |
pip wheel . --no-deps -w dist/
- uses: actions/upload-artifact@v4
with:
name: pure-wheel
path: dist/*.whl
build-cuda-linux:
name: Build CUDA wheel (Linux)
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: "3.12"
- name: Install torch (CUDA 12.8)
run: |
pip install torch --index-url https://download.pytorch.org/whl/cu128
- name: Setup CUDA
uses: Jimver/cuda-toolkit@v0.2.35
with:
cuda: "12.8.0"
- name: Build wheel (with CUDA kernels)
run: |
CSRC_KERNELS=true pip wheel . --no-deps --no-build-isolation -w dist/
- uses: actions/upload-artifact@v4
with:
name: cuda-wheel-linux
path: dist/*.whl
release:
name: Attach wheels to release
needs: [build-pure, build-cuda-linux]
runs-on: ubuntu-latest
permissions:
contents: write
steps:
- uses: actions/download-artifact@v4
with:
pattern: "*-wheel"
merge-multiple: true
- name: Create release & upload assets
uses: softprops/action-gh-release@v2
with:
files: ./*.whl
tag_name: ${{ github.ref_name }}
generate_release_notes: true
+15 -2
View File
@@ -5,8 +5,16 @@
!*/ !*/
# Allow specific file types and root files # Allow specific file types and root files
!*.py !astrai/**/*.py
!*.sh !scripts/**/*.py
!tests/**/*.py
!csrc/**/*.py
!csrc/**/*.cu
!csrc/**/*.h
!csrc/**/*.cuh
!scripts/**/*.sh
# Allow GitHub files # Allow GitHub files
!/.github/** !/.github/**
@@ -21,3 +29,8 @@
!/LICENSE !/LICENSE
!/pyproject.toml !/pyproject.toml
!/README.md !/README.md
# Allow extension modules (only source .py)
!/astrai/extension/**/*.py
# Allow build files
!/setup.py
+1 -1
View File
@@ -23,7 +23,7 @@ COPY astrai/ ./astrai/
COPY pyproject.toml . COPY pyproject.toml .
RUN pip install --no-cache-dir --upgrade pip \ RUN pip install --no-cache-dir --upgrade pip \
&& pip install --no-cache-dir . \ && pip install --no-cache-dir . \
--extra-index-url https://download.pytorch.org/whl/cu126 --extra-index-url https://download.pytorch.org/whl/cu128
# Production stage # Production stage
FROM ubuntu:24.04 AS production FROM ubuntu:24.04 AS production
+7 -8
View File
@@ -9,7 +9,7 @@
<div align="center"> <div align="center">
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python"> <img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
<img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license"> <img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license">
<img src="https://img.shields.io/github/v/release/ViperEkura/AstrAI?label=Release&color=76bad9" alt="release"> <img src="https://img.shields.io/github/v/tag/ViperEkura/AstrAI?label=Release&color=76bad9" alt="release">
<img src="https://img.shields.io/github/stars/ViperEkura/AstrAI?style=flat&label=Stars&color=76bad9" alt="stars"> <img src="https://img.shields.io/github/stars/ViperEkura/AstrAI?style=flat&label=Stars&color=76bad9" alt="stars">
<img src="https://img.shields.io/github/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks"> <img src="https://img.shields.io/github/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks">
</div> </div>
@@ -20,7 +20,7 @@
<a href="assets/docs/README-zh-CN.md">中文</a> • <a href="assets/docs/README-zh-CN.md">中文</a> •
<a href="https://github.com/ViperEkura/AstrAI/issues">Issue Tracker</a> • <a href="https://github.com/ViperEkura/AstrAI/issues">Issue Tracker</a> •
<a href="https://github.com/ViperEkura/AstrAI/discussions">Discussions</a> • <a href="https://github.com/ViperEkura/AstrAI/discussions">Discussions</a> •
<a href="https://huggingface.co/ViperEk/">HuggingFace</a> <a href="https://huggingface.co/ViperEkura">HuggingFace</a>
</div> </div>
<br> <br>
@@ -59,8 +59,9 @@ End-to-end walkthrough in 5 steps:
```bash ```bash
git clone https://github.com/ViperEkura/AstrAI.git git clone https://github.com/ViperEkura/AstrAI.git
cd AstrAI cd AstrAI
pip install -e . pip install -e . # pure PyTorch (no CUDA kernels)
# pip install -e ".[dev]" # optional: dev dependencies (pytest, ruff) # CSRC_KERNELS=true pip install -e . --no-build-isolation # optional: fused CUDA kernels
# pip install -e ".[dev]" # dev dependencies (pytest, ruff)
``` ```
**2. Download model** **2. Download model**
@@ -102,9 +103,7 @@ nohup python scripts/tools/train.py \
--warmup_ratio=0.05 \ --warmup_ratio=0.05 \
--max_lr=1e-4 \ --max_lr=1e-4 \
--max_grad_norm=1.0 \ --max_grad_norm=1.0 \
--adamw_beta1=0.9 \ --weight_decay=0.1 \
--adamw_beta2=0.95 \
--adamw_weight_decay=0.01 \
--window_size=2048 \ --window_size=2048 \
--ckpt_interval=10000 \ --ckpt_interval=10000 \
--ckpt_dir=./checkpoint \ --ckpt_dir=./checkpoint \
@@ -242,7 +241,7 @@ For major changes, please open an issue first to discuss what you would like to
- **GitHub Issues**: [Issue Tracker](https://github.com/ViperEkura/AstrAI/issues) - **GitHub Issues**: [Issue Tracker](https://github.com/ViperEkura/AstrAI/issues)
- **Discussions**: [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions) - **Discussions**: [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions)
- **HuggingFace**: [Model Hub](https://huggingface.co/ViperEk) - **HuggingFace**: [Model Hub](https://huggingface.co/ViperEkura)
### License ### License
+6 -7
View File
@@ -15,7 +15,7 @@
<div align="center"> <div align="center">
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python"> <img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
<img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license"> <img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license">
<img src="https://img.shields.io/github/v/release/ViperEkura/AstrAI?label=Release&color=76bad9" alt="release"> <img src="https://img.shields.io/github/v/tag/ViperEkura/AstrAI?label=Release&color=76bad9" alt="release">
<img src="https://img.shields.io/github/stars/ViperEkura/AstrAI?style=flat&label=Stars&color=76bad9" alt="stars"> <img src="https://img.shields.io/github/stars/ViperEkura/AstrAI?style=flat&label=Stars&color=76bad9" alt="stars">
<img src="https://img.shields.io/github/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks"> <img src="https://img.shields.io/github/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks">
</div> </div>
@@ -27,7 +27,7 @@
<a href="#chinese">中文</a> • <a href="#chinese">中文</a> •
<a href="https://github.com/ViperEkura/AstrAI/issues">问题追踪</a> • <a href="https://github.com/ViperEkura/AstrAI/issues">问题追踪</a> •
<a href="https://github.com/ViperEkura/AstrAI/discussions">讨论区</a> • <a href="https://github.com/ViperEkura/AstrAI/discussions">讨论区</a> •
<a href="https://huggingface.co/ViperEk">HuggingFace</a> <a href="https://huggingface.co/ViperEkura">HuggingFace</a>
</div> </div>
<br> <br>
@@ -65,7 +65,8 @@
```bash ```bash
git clone https://github.com/ViperEkura/AstrAI.git git clone https://github.com/ViperEkura/AstrAI.git
cd AstrAI cd AstrAI
pip install -e . pip install -e . # 纯 PyTorch(不含 CUDA 内核)
# CSRC_KERNELS=true pip install -e . --no-build-isolation # 可选:融合 CUDA 内核加速
# pip install -e ".[dev]" # 可选:开发依赖(pytest, ruff # pip install -e ".[dev]" # 可选:开发依赖(pytest, ruff
``` ```
@@ -108,9 +109,7 @@ nohup python scripts/tools/train.py \
--warmup_ratio=0.05 \ --warmup_ratio=0.05 \
--max_lr=1e-4 \ --max_lr=1e-4 \
--max_grad_norm=1.0 \ --max_grad_norm=1.0 \
--adamw_beta1=0.9 \ --weight_decay=0.1 \
--adamw_beta2=0.95 \
--adamw_weight_decay=0.01 \
--window_size=2048 \ --window_size=2048 \
--ckpt_interval=10000 \ --ckpt_interval=10000 \
--ckpt_dir=./checkpoint \ --ckpt_dir=./checkpoint \
@@ -248,7 +247,7 @@ SSE 流式格式、错误码和统计端点详见[推理文档](./inference.md)
- **GitHub Issues**: [问题追踪](https://github.com/ViperEkura/AstrAI/issues) - **GitHub Issues**: [问题追踪](https://github.com/ViperEkura/AstrAI/issues)
- **Discussions**: [GitHub 讨论区](https://github.com/ViperEkura/AstrAI/discussions) - **Discussions**: [GitHub 讨论区](https://github.com/ViperEkura/AstrAI/discussions)
- **HuggingFace**: [模型中心](https://huggingface.co/ViperEk) - **HuggingFace**: [模型中心](https://huggingface.co/ViperEkura)
### 许可证 ### 许可证
+154 -67
View File
@@ -21,6 +21,7 @@ classDiagram
class BaseModelConfig { class BaseModelConfig {
+Optional[str] model_type +Optional[str] model_type
+float neftune_alpha
+from_file(config_path) Self +from_file(config_path) Self
+to_file(config_path) +to_file(config_path)
} }
@@ -58,10 +59,11 @@ classDiagram
+Optional[int] dim_ffn +Optional[int] dim_ffn
+Optional[int] max_len +Optional[int] max_len
+Optional[float] rope_theta +Optional[float] rope_theta
+str attn_type
+Optional[int] n_heads +Optional[int] n_heads
+Optional[int] n_kv_heads +Optional[int] n_kv_heads
+Optional[bool] use_qk_norm +Optional[bool] use_qk_norm
+Optional[bool] use_gated_attention +str ffn_type
+Optional[dict] rope_scaling +Optional[dict] rope_scaling
+Optional[str] pooling_type +Optional[str] pooling_type
+Optional[bool] normalize_embeddings +Optional[bool] normalize_embeddings
@@ -115,14 +117,13 @@ classDiagram
+int n_epoch +int n_epoch
+int batch_per_device +int batch_per_device
+int grad_accum_steps +int grad_accum_steps
+float max_grad_norm +Optional[float] max_grad_norm
+list gradient_checkpointing_modules +list gradient_checkpointing_modules
+int start_epoch +int start_epoch
+int start_batch +int start_samples
+str ckpt_dir +str ckpt_dir
+int ckpt_interval +int ckpt_interval
+str log_dir +str log_dir
+int log_interval
+List[str] metrics +List[str] metrics
+Optional[LoRAConfig] lora +Optional[LoRAConfig] lora
+int random_seed +int random_seed
@@ -136,7 +137,9 @@ classDiagram
+str start_method +str start_method
+str device_type +str device_type
+Optional[Dataset] val_dataset +Optional[Dataset] val_dataset
+Optional[float] val_split
+int val_step +int val_step
+float neftune_alpha
+str parallel_mode +str parallel_mode
+dict executor_kwargs +dict executor_kwargs
+dict extra_kwargs +dict extra_kwargs
@@ -163,6 +166,13 @@ classDiagram
+__getitem__(index) Dict +__getitem__(index) Dict
} }
class RecordDataset {
+Optional[Callable] processor
+load(load_path, storage_type)
+__getitem__(index)
+__len__()
}
class DPODataset { class DPODataset {
+__getitem__(index) Dict +__getitem__(index) Dict
} }
@@ -174,13 +184,26 @@ classDiagram
class Store { class Store {
+Dict[str, List[Tensor]] _data +Dict[str, List[Tensor]] _data
+Dict[str, List[int]] _cum +Dict[str, List[int]] _cum
+Dict[str, List[int]] _offsets
+int _length +int _length
+int _num_records
+keys (property) +keys (property)
+load(path) +load(path)
+fetch(begin, end, keys)
+__len__() +__len__()
-_fetch_key(key, begin, end) Tensor -_normalize(raw, offsets)
-_normalize(raw) }
class Streamable {
<<mixin>>
+fetch(begin, end, keys)
-_fetch_stream_key(key, begin, end) Tensor
}
class Recordable {
<<mixin>>
+num_records (property)
+fetch_record(index, keys)
-_fetch_record_key(key, index) Tensor
} }
class H5Store { class H5Store {
@@ -192,6 +215,13 @@ classDiagram
+load(path) +load(path)
} }
class JsonlStore {
+JsonlSource _source
+Callable _processor
+load(path, transform, processor)
+fetch_record(index, keys)
}
class ResumableDistributedSampler { class ResumableDistributedSampler {
+int epoch +int epoch
+int iter +int iter
@@ -207,7 +237,7 @@ classDiagram
+Dict _entries +Dict _entries
+register(name) decorator +register(name) decorator
+create(train_type, window_size, stride) BaseDataset +create(train_type, window_size, stride) BaseDataset
+load(train_type, load_path, window_size, stride, storage_type) BaseDataset +load(train_type, load_path, window_size, stride, storage_type, tokenizer_path, max_len, store) BaseDataset
} }
} }
@@ -215,12 +245,13 @@ classDiagram
class Checkpoint { class Checkpoint {
+dict state_dict +dict state_dict
+int epoch +int epoch
+int iteration +int consumed_samples
+dict extra +dict extra
+dict meta +dict meta
+dict config +dict config
+save(save_dir) +save(save_dir)
+load(save_dir, broadcast) Checkpoint +load(save_dir, broadcast) Checkpoint
+load_any(save_dir, broadcast) Optional[Checkpoint]
} }
} }
@@ -350,7 +381,9 @@ classDiagram
class Embedding { class Embedding {
+Parameter weight +Parameter weight
+float neftune_noise_alpha
+forward(x) Tensor +forward(x) Tensor
+set_neftune_alpha(alpha)
} }
} }
@@ -372,6 +405,7 @@ classDiagram
+List[str] paths +List[str] paths
+str output_dir +str output_dir
+str tokenizer_path +str tokenizer_path
+AutoTokenizer tokenizer
+BaseMaskBuilder mask_builder +BaseMaskBuilder mask_builder
+PackingStrategy _packer +PackingStrategy _packer
+PositionIdStrategy _position_id +PositionIdStrategy _position_id
@@ -379,6 +413,18 @@ classDiagram
+transform(item) Optional[dict] +transform(item) Optional[dict]
+run() +run()
+_flush(domains, shard_idx) +_flush(domains, shard_idx)
+_inject_doc_reset_position_ids(keys, mode, seqs) Dict
+_inject_continuous_position_ids(tensors, mode, seqs) Dict
+_to_tensors(keys) Dict
}
class TokenizeTransform {
+PipelineConfig config
+AutoTokenizer tokenizer
+BaseMaskBuilder mask_builder
+PositionIdStrategy position_strategy
+from_config_file(path) TokenizeTransform
+apply(records) Dict[str, list]
} }
} }
@@ -407,7 +453,9 @@ classDiagram
+Dict _entries +Dict _entries
+register(name) decorator +register(name) decorator
+create(name, *args, **kwargs) T +create(name, *args, **kwargs) T
+get_component_class(name) Type
+list_registered() list +list_registered() list
+is_registered(name) bool
} }
class MaskBuilderFactory { class MaskBuilderFactory {
@@ -436,13 +484,15 @@ classDiagram
+dict model_config +dict model_config
+BaseExecutor executor +BaseExecutor executor
+int epoch +int epoch
+int iteration +int consumed_samples
+float loss +float loss
+float grad_norm
+DataLoader val_dataloader +DataLoader val_dataloader
+float val_loss +float val_loss
+int world_size +int world_size
+int rank +int rank
+dict kwargs +dict kwargs
+optimizer_step() int
} }
class TrainContextBuilder { class TrainContextBuilder {
@@ -485,14 +535,13 @@ classDiagram
} }
class GRPOStrategy { class GRPOStrategy {
+nn.Module old_model
+nn.Module ref_model +nn.Module ref_model
+float clip_eps +float clip_eps
+float kl_coef +float kl_coef
+int group_size +int group_size
+str reduction
+int sync_interval
+compute_loss(batch) Tensor +compute_loss(batch) Tensor
+sync_ref_model() +sync_old_model()
} }
class BaseScheduler { class BaseScheduler {
@@ -542,12 +591,12 @@ classDiagram
} }
class GradientClippingCallback { class GradientClippingCallback {
+float max_grad_norm +Optional[float] max_grad_norm
+on_optimizer_step(context) +on_optimizer_step(context)
} }
class GradientCheckpointingCallback { class GradientCheckpointingCallback {
+tuple modules +Optional[List[type]] modules
+on_train_begin(context) +on_train_begin(context)
+on_train_end(context) +on_train_end(context)
} }
@@ -561,31 +610,29 @@ classDiagram
+on_batch_end(context) +on_batch_end(context)
+on_train_end(context) +on_train_end(context)
+on_error(context) +on_error(context)
+save_extra(context) dict$ +save_extra(context) dict
} }
class ProgressBarCallback { class ProgressBarCallback {
+int num_epoch +int num_epoch
+int log_interval +int log_interval
+IO file +IO file
+tqdm progress_bar
+on_epoch_begin(context) +on_epoch_begin(context)
+on_batch_end(context) +on_optimizer_step(context)
+on_epoch_end(context) +on_epoch_end(context)
} }
class MetricLoggerCallback { class MetricCallback {
+Path log_dir +Path log_dir
+int save_interval +int save_interval
+int log_interval
+List[str] metrics +List[str] metrics
+on_batch_end(context) +int val_step
+on_optimizer_step(context)
+on_epoch_end(context)
+on_train_end(context) +on_train_end(context)
+on_error(context) +on_error(context)
}
class ValidationCallback {
-_run_validation(context) -_run_validation(context)
+on_optimizer_step(context)
} }
class CallbackFactory { class CallbackFactory {
@@ -594,18 +641,6 @@ classDiagram
+create(name, **kwargs) TrainCallback +create(name, **kwargs) TrainCallback
} }
class Muon {
+float lr
+float momentum
+float weight_decay
+bool nesterov
+int ns_steps
+Optional[float] adamw_lr
+tuple adamw_betas
+float adamw_eps
+float adamw_wd
+step(closure) Optional[float]
}
} }
namespace inference { namespace inference {
@@ -684,20 +719,44 @@ classDiagram
} }
class KVCache { class KVCache {
-PagePool _pool <<abstract>>
-Storage _storage
-TaskTable _table
+int page_size
+task_alloc(task_id, prompt_ids) bool +task_alloc(task_id, prompt_ids) bool
+task_free(task_id) +task_free(task_id)
+task_extend(task_id, pos) bool +task_extend(task_id, pos) bool
+task_cached(task_id) int +task_cached(task_id) int
+task_record_hashes(task_id, prompt_ids, start_logical_page) +task_record_hashes(task_id, prompt_ids, start_logical_page)
+make_table_tensor(task_ids, device) Tensor +bind_tasks(task_ids, total_len, device) CacheView
+bind(page_table, total_len) KvcacheView
} }
class KvcacheView { class PageCache {
+int page_size
-PagePool _pool
-Storage _storage
-TaskTable _table
+task_alloc(task_id, prompt_ids) bool
+task_free(task_id)
+task_extend(task_id, pos) bool
+task_cached(task_id) int
+task_record_hashes(task_id, prompt_ids, start_logical_page)
+bind_tasks(task_ids, total_len, device) PageCacheView
}
class ContiguousCache {
+int max_seq_len
+Tensor k, v
+task_alloc(task_id, prompt_ids) bool
+task_free(task_id)
+task_extend(task_id, pos) bool
+bind_tasks(task_ids, total_len, device) ContiguousCacheView
}
class CacheView {
<<abstract>>
+write(layer_id, k, v)
+gather(layer_id) Tuple[Tensor, Tensor]
}
class PageCacheView {
-Storage _storage -Storage _storage
+Tensor _page_table +Tensor _page_table
+int _total_len +int _total_len
@@ -705,6 +764,14 @@ classDiagram
+gather(layer_id) Tuple[Tensor, Tensor] +gather(layer_id) Tuple[Tensor, Tensor]
} }
class ContiguousCacheView {
-ContiguousCache _cache
+Tensor _batch_indices
+int _total_len
+write(layer_id, k, v)
+gather(layer_id) Tuple[Tensor, Tensor]
}
class TaskTable { class TaskTable {
+set(task_id, page_table, cached) +set(task_id, page_table, cached)
+get(task_id) List[int] +get(task_id) List[int]
@@ -727,7 +794,6 @@ classDiagram
+int output_tokens +int output_tokens
+float arrival_time +float arrival_time
+Optional[float] finish_time +Optional[float] finish_time
+Optional[Callable] stream_callback
+int next_pos +int next_pos
+is_finished(stop_ids) bool +is_finished(stop_ids) bool
} }
@@ -810,7 +876,9 @@ classDiagram
class ChatMessage { class ChatMessage {
+str role +str role
+str content +Optional[str] content
+Optional[List[Dict]] tool_calls
+Optional[str] tool_call_id
} }
class ChatCompletionRequest { class ChatCompletionRequest {
@@ -827,6 +895,8 @@ classDiagram
+Optional[float] frequency_penalty +Optional[float] frequency_penalty
+Optional[Dict[int, float]] logit_bias +Optional[Dict[int, float]] logit_bias
+Optional[str] user +Optional[str] user
+Optional[List[ToolDef]] tools
+Optional[Union[str, Dict]] tool_choice
} }
class AnthropicMessage { class AnthropicMessage {
@@ -850,7 +920,7 @@ classDiagram
<<abstract>> <<abstract>>
+prepare(request, engine) Tuple[str, GenContext, List[str]] +prepare(request, engine) Tuple[str, GenContext, List[str]]
+format_stream_start(ctx) List[str] +format_stream_start(ctx) List[str]
+format_chunk(token) str +format_chunk(token) List[str]
+format_stream_end(ctx, stop) List[str] +format_stream_end(ctx, stop) List[str]
+format_response(ctx, content, stop) Dict +format_response(ctx, content, stop) Dict
} }
@@ -858,7 +928,7 @@ classDiagram
class OpenAIResponseBuilder { class OpenAIResponseBuilder {
+prepare(request, engine) Tuple +prepare(request, engine) Tuple
+format_stream_start(ctx) List[str] +format_stream_start(ctx) List[str]
+format_chunk(token) str +format_chunk(token) List[str]
+format_stream_end(ctx, stop) List[str] +format_stream_end(ctx, stop) List[str]
+format_response(ctx, content, stop) Dict +format_response(ctx, content, stop) Dict
} }
@@ -866,7 +936,7 @@ classDiagram
class AnthropicResponseBuilder { class AnthropicResponseBuilder {
+prepare(request, engine) Tuple +prepare(request, engine) Tuple
+format_stream_start(ctx) List[str] +format_stream_start(ctx) List[str]
+format_chunk(token) str +format_chunk(token) List[str]
+format_stream_end(ctx, stop) List[str] +format_stream_end(ctx, stop) List[str]
+format_response(ctx, content, stop) Dict +format_response(ctx, content, stop) Dict
} }
@@ -1031,14 +1101,21 @@ classDiagram
TrainCallback <|-- GradientCheckpointingCallback TrainCallback <|-- GradientCheckpointingCallback
TrainCallback <|-- CheckpointCallback TrainCallback <|-- CheckpointCallback
TrainCallback <|-- ProgressBarCallback TrainCallback <|-- ProgressBarCallback
TrainCallback <|-- MetricLoggerCallback TrainCallback <|-- MetricCallback
TrainCallback <|-- ValidationCallback
BaseDataset <|-- SEQDataset BaseDataset <|-- SEQDataset
BaseDataset <|-- SFTDataset BaseDataset <|-- SFTDataset
BaseDataset <|-- DPODataset BaseDataset <|-- RecordDataset
BaseDataset <|-- GRPODataset RecordDataset <|-- DPODataset
RecordDataset <|-- GRPODataset
Store <|-- H5Store Store <|-- H5Store
Store <|-- MmapStore Store <|-- MmapStore
Store <|-- JsonlStore
H5Store --|> Streamable
H5Store --|> Recordable
MmapStore --|> Streamable
MmapStore --|> Recordable
JsonlStore --|> Streamable
JsonlStore --|> Recordable
BaseSamplingStrategy <|-- TemperatureStrategy BaseSamplingStrategy <|-- TemperatureStrategy
BaseSamplingStrategy <|-- TopKStrategy BaseSamplingStrategy <|-- TopKStrategy
BaseSamplingStrategy <|-- TopPStrategy BaseSamplingStrategy <|-- TopPStrategy
@@ -1071,11 +1148,15 @@ classDiagram
ResponseBuilder <|-- OpenAIResponseBuilder ResponseBuilder <|-- OpenAIResponseBuilder
ResponseBuilder <|-- AnthropicResponseBuilder ResponseBuilder <|-- AnthropicResponseBuilder
BaseMaskBuilder <|-- SectionedMaskBuilder BaseMaskBuilder <|-- SectionedMaskBuilder
KVCache <|-- PageCache
KVCache <|-- ContiguousCache
CacheView <|-- PageCacheView
CacheView <|-- ContiguousCacheView
%% --- Composition (strong ownership, part destroyed with whole) --- %% --- Composition (strong ownership, part destroyed with whole) ---
KVCache *-- PagePool PageCache *-- PagePool
KVCache *-- Storage PageCache *-- Storage
KVCache *-- TaskTable PageCache *-- TaskTable
InferenceEngine *-- InferenceScheduler InferenceEngine *-- InferenceScheduler
InferenceScheduler *-- KVCache InferenceScheduler *-- KVCache
InferenceScheduler *-- Executor InferenceScheduler *-- Executor
@@ -1103,11 +1184,15 @@ classDiagram
TrainContext o-- BaseScheduler TrainContext o-- BaseScheduler
TrainContext o-- Checkpoint TrainContext o-- Checkpoint
TrainContext o-- BaseExecutor TrainContext o-- BaseExecutor
KvcacheView o-- Storage PageCacheView o-- Storage
ContiguousCacheView o-- ContiguousCache
SamplingPipeline o-- BaseSamplingStrategy SamplingPipeline o-- BaseSamplingStrategy
BaseDataset o-- Store BaseDataset o-- Store
Pipeline o-- PipelineConfig Pipeline o-- PipelineConfig
Pipeline o-- BaseMaskBuilder Pipeline o-- BaseMaskBuilder
Pipeline o-- AutoTokenizer
TokenizeTransform o-- AutoTokenizer
TokenizeTransform o-- BaseMaskBuilder
%% --- Dependency (uses temporarily) --- %% --- Dependency (uses temporarily) ---
TrainConfig ..> BaseStrategy : selects TrainConfig ..> BaseStrategy : selects
@@ -1125,6 +1210,7 @@ classDiagram
DecoderBlock ..> FFNFactory : uses DecoderBlock ..> FFNFactory : uses
StoreFactory ..> H5Store : creates StoreFactory ..> H5Store : creates
StoreFactory ..> MmapStore : creates StoreFactory ..> MmapStore : creates
StoreFactory ..> JsonlStore : creates
ConfigFactory ..> AutoRegressiveLMConfig : creates ConfigFactory ..> AutoRegressiveLMConfig : creates
ConfigFactory ..> EncoderConfig : creates ConfigFactory ..> EncoderConfig : creates
ExecutorFactory ..> NoneExecutor : creates ExecutorFactory ..> NoneExecutor : creates
@@ -1138,7 +1224,8 @@ classDiagram
TrainContextBuilder ..> ResumableDistributedSampler : creates TrainContextBuilder ..> ResumableDistributedSampler : creates
Checkpoint ..> Checkpoint : serializes Checkpoint ..> Checkpoint : serializes
CheckpointCallback ..> Checkpoint : creates CheckpointCallback ..> Checkpoint : creates
KVCache ..> KvcacheView : binds PageCache ..> PageCacheView : binds
ContiguousCache ..> ContiguousCacheView : binds
InferenceEngine ..> GenerationRequest : uses InferenceEngine ..> GenerationRequest : uses
InferenceEngine ..> GenerateResult : creates InferenceEngine ..> GenerateResult : creates
OpenAIResponseBuilder ..> ChatCompletionRequest : receives OpenAIResponseBuilder ..> ChatCompletionRequest : receives
@@ -1149,7 +1236,7 @@ classDiagram
%% --- Association (general usage) --- %% --- Association (general usage) ---
Trainer --> TrainConfig Trainer --> TrainConfig
DPOStrategy --> AutoModel DPOStrategy --> AutoModel
GRPOStrategy --> AutoModel GRPOStrategy --> AutoModel : policy/old/ref
InferenceScheduler --> Task InferenceScheduler --> Task
InferenceScheduler --> TaskStatus InferenceScheduler --> TaskStatus
Task --> TaskStatus Task --> TaskStatus
@@ -1166,22 +1253,22 @@ classDiagram
| Module | Components | Description | | 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.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, Pipeline, filter_by_length, PackingStrategy, PackingStrategyFactory, PositionIdStrategy, PositionIdStrategyFactory, StoreWriter, StoreWriterFactory | Declarative JSON-driven data preprocessing | | **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** | BaseDatasetGRPODataset, StoreMmapStore, StoreFactory, ResumableDistributedSampler, DatasetFactory | Dataset loading and management | | **astrai.dataset** | BaseDatasetRecordDatasetDPO/GRPODataset, SEQDataset, SFTDataset, Store, Streamable, Recordable, H5Store, MmapStore, JsonlStore, StoreFactory, ResumableDistributedSampler, DatasetFactory | Dataset loading and management |
| **astrai.serialization** | Checkpoint | Model serialization | | **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, RotaryEmbedding, Embedding | Neural network model |
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template | | **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategyGRPOStrategy, StrategyFactory, BaseSchedulerWSDScheduler, SchedulerFactory, TrainCallback(Protocol)ValidationCallback, CallbackFactory, Muon | Training workflow | | **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategyGRPOStrategy, StrategyFactory, BaseSchedulerWSDScheduler, SchedulerFactory, TrainCallback(Protocol)MetricCallback, CallbackFactory | Training workflow |
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, KVCacheKvcacheView, AllocatorStorage, Task, TaskManager, TaskStatus, GenerationRequest, GenerateResult, BaseSamplingStrategySamplingPipeline, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, ChatMessageMessagesRequest, app | Inference service | | **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, KVCacheContiguousCache/PageCache, CacheViewContiguousCacheView/PageCacheView, AllocatorStorage, Task, TaskManager, TaskStatus, GenerationRequest, GenerateResult, BaseSamplingStrategySamplingPipeline, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, ChatMessageMessagesRequest, 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.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.factory** | Registry, BaseFactory[T] | Component registration | | **astrai.factory** | BaseFactory | Component registration |
| **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers | | **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers |
## Design Patterns ## Design Patterns
| Pattern | Classes | Purpose | | Pattern | Classes | Purpose |
|---------|---------|---------| |---------|---------|---------|
| **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory` | Decorator-based component creation | | **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory`, `MaskBuilderFactory`, `StoreWriterFactory`, `PackingStrategyFactory`, `PositionIdStrategyFactory` | Decorator-based component creation |
| **Registry** | `BaseFactory` | Component registration | | **Registry** | `BaseFactory` | Component registration |
| **Strategy** | `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy` | Training strategy switching | | **Strategy** | `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy` | Training strategy switching |
| **Strategy (Sampling)** | `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations | | **Strategy (Sampling)** | `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations |
@@ -1191,7 +1278,7 @@ classDiagram
| **Context** | `TrainContext` | Unified training state bag | | **Context** | `TrainContext` | Unified training state bag |
| **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction | | **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` | Gradient accumulation & model distribution |
| **Storage** | `Store`, `H5Store`, `MmapStore` | Format-agnostic data access with multi-segment support | | **Storage** | `Store`, `H5Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support |
| **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching | | **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching |
| **AutoModel Registry** | `AutoModel`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading | | **AutoModel Registry** | `AutoModel`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading |
@@ -1203,10 +1290,10 @@ classDiagram
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`
5. **Inference Flow**: `InferenceEngine``InferenceScheduler``AutoRegressiveLM`, backed by `KVCache` + `SamplingPipeline` 5. **Inference Flow**: `InferenceEngine``InferenceScheduler``AutoRegressiveLM`, backed by `KVCache` + `SamplingPipeline`
6. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP 6. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP
7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (H5Store/MmapStore) loads data with explicit `_length` and multi-segment `_data` 7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (H5Store/MmapStore/JsonlStore) loads data with explicit `_length` and multi-segment `_data`
8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata (rank-0 only), extra state saved as `{key}.pt` 8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata (rank-0 only), extra state saved as `{key}.pt`
9. **Scheduler**: `SchedulerFactory` creates `CosineScheduler`/`SGDRScheduler`/`WSDScheduler` 9. **Scheduler**: `SchedulerFactory` creates `CosineScheduler`/`SGDRScheduler`/`WSDScheduler`
10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops 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 11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers
> Document Update Time: 2026-05-30 > Document Update Time: 2026-07-19
+40 -19
View File
@@ -46,10 +46,10 @@ The output `meta.json` records the storage format, key names, dtype, total token
### Format Detection ### Format Detection
`detect_format(load_path)` inspects the directory: `detect_format(load_path)` inspects the path:
- If `*.h5` files exist → `"h5"` (HDF5 backend) - If `load_path` is a file: checks suffix — `.h5`/`.hdf5``"h5"`, `.jsonl``"jsonl"`, unknown suffix raises `ValueError`
- If `*.bin` + `meta.json` files exist → `"bin"` (memory-mapped backend) - If `load_path` is a directory: recursively globs for `*.h5`/`*.hdf5` files → `"h5"`, `*.bin` + `**/meta.json` `"bin"`, or `*.jsonl` + `dataset_config.json``"jsonl"`
### Store Backends ### Store Backends
@@ -58,41 +58,62 @@ Storage format is auto-detected by `detect_format()`; backends are dispatched vi
``` ```
StoreFactory.create("h5") → H5Store StoreFactory.create("h5") → H5Store
StoreFactory.create("bin") → MmapStore StoreFactory.create("bin") → MmapStore
StoreFactory.create("jsonl") → JsonlStore
``` ```
**H5Store**: Reads HDF5 files, supports `share_memory_()` for multi-process DataLoader workers (copies tensors to shared memory). All three inherit `Store` (base, owns `_data`/`_cum`/`_offsets`/`_normalize`) plus the `Streamable` and `Recordable` mixins, so every backend supports both `fetch(begin, end, keys)` (stream) and `fetch_record(index, keys)` (record) APIs.
**MmapStore**: Memory-maps `.bin` files. OS page cache sharing is native — no explicit `share_memory_()` needed. Uses `torch.from_numpy(np.memmap(...))`. **H5Store**: Reads HDF5 files. Tensors are loaded into host memory and normalized into segmented storage. `segments_are_records=True` — each `data_i` dataset is one record.
Both backends normalise tensors into `Store._data[Dict[str, List[Tensor]]]` + `Store._cum[Dict[str, List[int]]]` (cumulative lengths for bisect-based indexing). **MmapStore**: Memory-maps `.bin` files. OS page cache sharing is native — no explicit `share_memory_()` needed. Uses `torch.from_numpy(np.memmap(...))`. `segments_are_records=False` — bin segments are contiguous streams; record access is driven by `_offsets` (written when `save_bin(..., record_keys=...)` was used at preprocessing time).
**JsonlStore**: On-the-fly tokenization of raw JSONL files at load time. Requires a `dataset_config.json` alongside the `.jsonl` files following the same `PipelineConfig` schema with an additional `tokenizer_path` field. Two modes: eager (default, applies `TokenizeTransform` to all records at load) and lazy (`processor=fn` given, defers tokenisation to `fetch_record` — used by DPO/GRPO).
All backends normalise tensors into `Store._data[Dict[str, List[Tensor]]]` + `Store._cum[Dict[str, List[int]]]` (cumulative lengths for bisect-based stream indexing) + `Store._offsets[Dict[str, List[int]]]` (per-record offsets for record-mode indexing). Nested keys (GRPO `responses`/`masks` as `List[List[Tensor]]`) are stored as-is and excluded from both bookkeepings — they are only accessed record-by-record.
## Data Keys by Training Type ## Data Keys by Training Type
| Type | Storage Keys | | Type | Storage Keys | Access Mode |
|------|-------------| |------|-------------|-------------|
| `seq` | `sequence` (→ input_ids, target_ids via offset-by-1) | | `seq` | `sequence` (→ input_ids, target_ids via offset-by-1) | stream (`fetch`) |
| `sft` | `sequence`, `loss_mask`, `position_ids` | | `sft` | `sequence`, `loss_mask`, `position_ids` | stream (`fetch`) |
| `dpo` | `chosen`, `rejected`, `chosen_mask`, `rejected_mask` | | `dpo` | `chosen`, `rejected`, `chosen_mask`, `rejected_mask` | record (`fetch_record`) |
| `grpo` | `prompts`, `responses`, `masks`, `rewards` | | `grpo` | `prompts`, `responses`, `masks`, `rewards` | record (`fetch_record`) |
## Dataset Architecture ## Dataset Architecture
``` ```
DatasetFactory.load(train_type, load_path, window_size, stride=None, storage_type=None) DatasetFactory.load(train_type, load_path, window_size, stride=None,
storage_type=None, tokenizer_path=None,
max_len=2048, store=None)
→ BaseDataset.load(load_path, storage_type=None) → BaseDataset.load(load_path, storage_type=None)
→ detect_format(load_path) → detect_format(load_path)
→ StoreFactory.create(storage_type) → StoreFactory.create(storage_type)
→ Store.load(load_path) → Store.load(load_path)
H5Store._normalize() / MmapStore._normalize() → _normalize(raw) # base Store, shared by both backends
→ Store._data[Dict[str, List[Tensor]]] + _cum[Dict[str, List[int]]] → Store._data[Dict[str, List[Tensor]]]
→ BaseDataset.__getitem__(idx) + _cum[Dict[str, List[int]]] (stream mode)
+ _offsets[Dict[str, List[int]]] (record mode)
Stream datasets (SEQ/SFT):
BaseDataset.__getitem__(idx)
→ get_index(idx) → [begin, end) → get_index(idx) → [begin, end)
→ Store.fetch(begin, end, keys) → Tensor / Dict[str, Tensor] → Store.fetch(begin, end, keys) → Tensor / Dict[str, Tensor]
Record datasets (DPO/GRPO via RecordDataset):
RecordDataset.__getitem__(idx)
→ Store.fetch_record(idx, keys) → Tensor / Dict[str, Tensor]
``` ```
`window_size` = max input length, `stride` = step between consecutive samples (defaults to `window_size`, optional). `storage_type` defaults to `None` (auto-detect via `detect_format`). Class hierarchy: `BaseDataset``SEQDataset` / `SFTDataset` (stream); `BaseDataset``RecordDataset``DPODataset` / `GRPODataset` (record).
`Store.fetch(begin, end, keys)` accepts a single key (`str`) returning a `Tensor`, or a list of keys returning `Dict[str, Tensor]`. Internally uses `bisect` across multi-segment tensors. Raises `RuntimeError("Store not loaded")` if called before `load()`. `window_size` = max input length, `stride` = step between consecutive samples (defaults to `window_size`, optional). Only meaningful for stream datasets — record datasets ignore both. `storage_type` defaults to `None` (auto-detect via `detect_format`).
`tokenizer_path` triggers lazy on-the-fly tokenisation for record datasets on raw JSONL (DPO builds a `dpo_processor`; SEQ/SFT/pre-tokenised backends ignore it). `store` (pre-built `Store`) bypasses `load_path`/`storage_type`/`tokenizer_path` entirely — the caller controls Store construction.
`Store.fetch(begin, end, keys)` (stream mode, on `Streamable`): accepts a single key (`str`) returning a `Tensor`, or a list of keys returning `Dict[str, Tensor]`. Internally uses `bisect` across multi-segment tensors. Raises `RuntimeError("Store not loaded")` if called before `load()`.
`Store.fetch_record(index, keys)` (record mode, on `Recordable`): same key API. Uses `_offsets[key]` when present (bin layout with per-record offsets), otherwise indexes `_data[key]` directly (H5/JSONL where each segment is one record).
## Sampler ## Sampler
@@ -106,4 +127,4 @@ DatasetFactory.load(train_type, load_path, window_size, stride=None, storage_typ
Standard PyTorch `DataLoader` with configurable `batch_size`, `num_workers`, `pin_memory`, `prefetch_factor`. Sampler produces indices; dataloader fetches tensor batches via `__getitem__`. Standard PyTorch `DataLoader` with configurable `batch_size`, `num_workers`, `pin_memory`, `prefetch_factor`. Sampler produces indices; dataloader fetches tensor batches via `__getitem__`.
> Document Update Time: 2026-06-19 > Document Update Time: 2026-07-19
+36 -33
View File
@@ -23,29 +23,40 @@ RoPE is applied **before** KV cache write, not after — otherwise position enco
## KVCache System ## KVCache System
Six classes (plus two helpers) working together: Seven classes working together, with two concrete cache implementations:
### ContiguousCache (default)
``` ```
KVCache (facade) 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.
### PageCache (paged with prefix sharing)
```
PageCache (paged KV cache with prefix sharing, alternative)
├── PagePool orchestrates page allocation + prefix matching ├── PagePool orchestrates page allocation + prefix matching
│ ├── Allocator bitmask-based page allocator + ref-count + LRU eviction (inside PagePool) │ ├── Allocator bitmask-based page allocator + ref-count + LRU
│ └── PrefixCache hash-based prefix matching (page_hash via polynomial hash) (inside PagePool) │ └── PrefixCache hash-based prefix matching (page_hash via polynomial hash)
├── TaskTable maps task_id → page_table + cached token count ├── 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 (n_layers × n_pages × page_size × n_kv_heads × head_dim)
└── KvcacheView bundles Storage + page_table + total_len for attention layers (returned by bind()) └── PageCacheView bundles Storage + page_table + total_len for attention layers
``` ```
`KVCache.bind(page_table, total_len)` returns a `KvcacheView` used by attention layers via `write()` / `gather()`. `isinstance(cache, KVCache)` checks dispatch to the correct view. Both implement the abstract `KVCache` interface used by `Executor` and `InferenceScheduler`.
## Continuous Batching ## Continuous Batching
`InferenceScheduler` runs a daemon thread with a 4-phase loop: `InferenceScheduler` runs a daemon thread with a 4-phase loop:
``` ```
1. Cleanup → Remove finished tasks, free KV pages 1. Cleanup → Remove finished tasks, free KV cache slots/pages
2. Refill → Pop from waiting_queue, task_alloc pages, activate 2. Refill → Pop from waiting_queue, task_alloc resources, activate
3. Prefill → Group by (prompt_len, start_pos), run full forward 3. Prefill → Group by (prompt_len, start_pos), run full forward
4. Decode → Pick largest same-position group, single-token forward 4. Decode → Run single-token forward for each same-position group
``` ```
## Sampling (Strategy Pattern) ## Sampling (Strategy Pattern)
@@ -152,12 +163,13 @@ Supports `stop_sequences` and streaming via `event: content_block_delta`.
data: {"id":"chatcmpl-...","object":"chat.completion.chunk","created":...,"model":"astrai", data: {"id":"chatcmpl-...","object":"chat.completion.chunk","created":...,"model":"astrai",
"choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]} "choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}
data: {"id":"chatcmpl-...","object":"chat.completion.chunk",..., data: {"id":"chatcmpl-...","object":"chat.completion.chunk","created":0,"model":"astrai",
"choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]} "choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]}
data: {"id":"chatcmpl-...","object":"chat.completion.chunk",..., data: {"id":"chatcmpl-...","object":"chat.completion.chunk","created":...,"model":"astrai",
"choices":[{"index":0,"delta":{},"finish_reason":"stop"}], "choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}
"usage":{"prompt_tokens":5,"completion_tokens":1,"total_tokens":6}}
data: {"prompt_tokens":5,"completion_tokens":1,"total_tokens":6}
data: [DONE] data: [DONE]
``` ```
@@ -167,7 +179,7 @@ data: [DONE]
``` ```
event: message_start event: message_start
data: {"type":"message_start","message":{"id":"msg_...","model":"astrai","role":"assistant", data: {"type":"message_start","message":{"id":"msg_...","model":"astrai","role":"assistant",
"content":[],"stop_reason":null,...}} "content":[],"usage":{"input_tokens":0}}}
event: content_block_start event: content_block_start
data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}} data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}
@@ -179,7 +191,7 @@ event: content_block_stop
data: {"type":"content_block_stop","index":0} data: {"type":"content_block_stop","index":0}
event: message_delta event: message_delta
data: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{...}} data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},"usage":{...}}
event: message_stop event: message_stop
data: {"type":"message_stop"} data: {"type":"message_stop"}
@@ -187,26 +199,20 @@ data: {"type":"message_stop"}
### Error Responses ### Error Responses
All endpoints use standard HTTP status codes: The server returns standard HTTP status codes. Pydantic validation errors (e.g. missing required fields)
are handled automatically by FastAPI with 422 status. The only application-level error is engine initialization:
| Status | Meaning | | Status | Meaning |
|--------|---------| |--------|---------|
| 200 | Success | | 200 | Success |
| 400 | Invalid request (bad JSON, missing fields, validation error) |
| 405 | Method not allowed |
| 422 | Unprocessable entity (Pydantic validation) | | 422 | Unprocessable entity (Pydantic validation) |
| 500 | Internal server error (model crash, OOM, scheduler failure) |
| 503 | Service unavailable (model not loaded, engine not ready) | | 503 | Service unavailable (model not loaded, engine not ready) |
Error response body: Error response body (503):
```json ```json
{ {
"error": { "detail": "Engine not initialized"
"message": "Invalid request: max_tokens must be > 0",
"type": "invalid_request_error",
"code": 400
}
} }
``` ```
@@ -220,16 +226,13 @@ Response:
```json ```json
{ {
"active_requests": 3, "total_tasks": 128,
"waiting_requests": 2, "total_tokens": 10240,
"total_requests": 128, "active_tasks": 3,
"cache_usage": 0.45, "waiting_queue": 2
"tokens_generated": 10240
} }
``` ```
`cache_usage` is the fraction of KV cache pages currently in use (0.01.0).
## Engine API ## Engine API
```python ```python
@@ -246,4 +249,4 @@ async for token in engine.generate_async("Hello", ...): # -> AsyncGenerator[s
print(token) print(token)
``` ```
> Document Update Time: 2026-06-19 > Document Update Time: 2026-07-09
+26 -14
View File
@@ -26,15 +26,19 @@
|-----------|-------------|---------| |-----------|-------------|---------|
| `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 | | `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 |
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 | | `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
| `--max_grad_norm` | Maximum gradient norm for clipping | 1.0 | | `--max_grad_norm` | Maximum gradient norm for clipping (None disables) | None |
### Optimizer (AdamW) ### Optimizer (MuonMix)
Combined optimizer: matrix parameters via **Muon**, non-matrix via **AdamW** (`fused=True`).
| Parameter | Description | Default | | Parameter | Description | Default |
|-----------|-------------|---------| |-----------|-------------|---------|
| `--adamw_beta1` | AdamW beta1 | 0.9 | | `--weight_decay` | Weight decay (applied to Muon matrix params; non-matrix use 0) | 0.1 |
| `--adamw_beta2` | AdamW beta2 | 0.95 | | `--muon_momentum` | Muon momentum factor | 0.95 |
| `--adamw_weight_decay` | AdamW weight decay | 0.01 | | `--muon_nesterov` | Enable Nesterov momentum for Muon | True |
| `--muon_ns_steps` | Newton-Schulz iteration steps for Muon | 5 |
| `--muon_adjust_lr` | Muon LR adjustment strategy (`original`, `match_rms_adamw`) | `match_rms_adamw` |
### Data Loading ### Data Loading
@@ -53,7 +57,7 @@
| `--ckpt_interval` | Iterations between checkpoints | 5000 | | `--ckpt_interval` | Iterations between checkpoints | 5000 |
| `--ckpt_dir` | Checkpoint save directory | checkpoint | | `--ckpt_dir` | Checkpoint save directory | checkpoint |
| `--start_epoch` | Resume from epoch (0 = from scratch) | 0 | | `--start_epoch` | Resume from epoch (0 = from scratch) | 0 |
| `--start_batch` | Resume from batch iteration | 0 | | `--start_samples` | Resume from sample count per rank | 0 |
### Validation ### Validation
@@ -67,8 +71,7 @@
| Parameter | Description | Default | | Parameter | Description | Default |
|-----------|-------------|---------| |-----------|-------------|---------|
| `--log_dir` | Directory for metric logs | checkpoint/logs | | `--log_dir` | Directory for metric logs | checkpoint/logs |
| `--log_interval` | Number of batch iterations between metric logs | 100 | | `--metrics` | Metrics to log (e.g. --metrics loss lr val_loss) | ["loss", "lr", "grad_norm"] |
| `--metrics` | Metrics to log (e.g. --metrics loss lr val_loss) | ["loss", "lr"] |
### Gradient Checkpointing ### Gradient Checkpointing
@@ -100,6 +103,17 @@
| `--grpo_sync_interval` | GRPO ref_model sync interval (steps) | 200 | `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` | | `--neftune_alpha` | NEFTune noise alpha (0=disabled, typical: 5.0) | 0.0 | `sft` |
### 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) |
| `--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) |
| `--decay_steps` | WSD decay steps | None (total_steps - warmup_steps - stable_steps) |
### Usage Example ### Usage Example
```bash ```bash
@@ -116,9 +130,7 @@ nohup python scripts/tools/train.py \
--warmup_ratio=0.05 \ --warmup_ratio=0.05 \
--max_lr=1e-4 \ --max_lr=1e-4 \
--max_grad_norm=1.0 \ --max_grad_norm=1.0 \
--adamw_beta1=0.9 \ --weight_decay=0.1 \
--adamw_beta2=0.95 \
--adamw_weight_decay=0.01 \
--window_size=2048 \ --window_size=2048 \
--ckpt_interval=10000 \ --ckpt_interval=10000 \
--ckpt_dir=./checkpoint \ --ckpt_dir=./checkpoint \
@@ -161,7 +173,7 @@ See [Inference Guide](inference.md) for HTTP API documentation.
| `--top_k` | int | `30` | Top-k filtering | | `--top_k` | int | `30` | Top-k filtering |
| `--top_p` | float | `0.95` | Nucleus sampling threshold | | `--top_p` | float | `0.95` | Nucleus sampling threshold |
| `--batch_size` | int | `1` | Batch size for generation | | `--batch_size` | int | `1` | Batch size for generation |
| `--max_tokens` | int | `2048` | Maximum tokens to generate | | `--max_tokens` | int | model config `max_len` | Maximum tokens to generate |
Usage: Usage:
```bash ```bash
@@ -178,7 +190,7 @@ python scripts/tools/generate.py \
| `input_files` | path(s) | required | Input JSONL file(s), supports glob (`data/*.jsonl`) | | `input_files` | path(s) | required | Input JSONL file(s), supports glob (`data/*.jsonl`) |
| `--output_dir`, `-o` | path | required | Output directory for processed data | | `--output_dir`, `-o` | path | required | Output directory for processed data |
| `--config`, `-c` | path | required | Preprocessing pipeline config (JSON) | | `--config`, `-c` | path | required | Preprocessing pipeline config (JSON) |
| `--num_workers` | int | `4` | Number of parallel workers | | `--tokenizer_path` | str | `params` | Path to tokenizer directory |
Usage: Usage:
```bash ```bash
@@ -189,4 +201,4 @@ See [Preprocessing Guide](preprocessing.md) for config file format and examples.
--- ---
> Document Update Time: 2026-06-19 > Document Update Time: 2026-07-19
+16 -13
View File
@@ -1,6 +1,6 @@
# Preprocessing Pipeline # Preprocessing Pipeline
Declarative JSON-driven data preprocessing. One `SectionedMaskBuilder` handles all formats via `input.sections` (single-output) or `input.sources` (multi-output). Declarative JSON-driven data preprocessing. `MaskBuilderFactory` supports three registered builders: `"single"` (single-output via `input.sections`), `"multi"` (multi-output via `input.sources`), and `"sectioned"` (façade dispatching to `single` or `multi` based on config).
## Contents ## Contents
@@ -26,8 +26,9 @@ A single config file captures the entire pipeline, reusable and version-controll
```json ```json
{ {
"version": 1,
"input": {}, // sections (single) or sources (multi) "input": {}, // sections (single) or sources (multi)
"mask": {}, // role "train" | "mask" "mask": {}, // role -> "train" | "mask"
"mask_default": "mask", "mask_default": "mask",
"preprocessing": {}, "preprocessing": {},
"output": {} "output": {}
@@ -220,11 +221,12 @@ Config:
} }
``` ```
Output keys: `prompts`, `responses`, `masks`, `rewards` (float32) Output keys: `prompts`, `prompts_mask`, `responses`, `masks`, `rewards` (float32)
- `action: "value"` — extract raw values from JSONL without tokenisation - `action: "value"` — extract raw values from JSONL without tokenisation
- `list_field: true` — tokenise each list element independently, then concatenate - `list_field: true` — tokenise each list element independently, then concatenate
- `mask_key: "masks"` — rename the auto-generated mask key (default: `responses_mask`) - `mask_key: "masks"` — rename the auto-generated mask key (default: `responses_mask`)
- `prompts_mask` is auto-generated (all masked) and unused by GRPOStrategy
--- ---
@@ -266,7 +268,7 @@ When `sources` is set, `sections` is ignored.
| `storage_format` | str | `"bin"` | `"bin"` (mmap) or `"h5"` | | `storage_format` | str | `"bin"` | `"bin"` (mmap) or `"h5"` |
| `max_tokens_per_shard` | int | `100000000` | Flush threshold in cumulative tokens | | `max_tokens_per_shard` | int | `100000000` | Flush threshold in cumulative tokens |
| `dtype` | dict[str, str] | `{}` | Per-key tensor dtype override (e.g. `{"loss_mask": "bool"}`) | | `dtype` | dict[str, str] | `{}` | Per-key tensor dtype override (e.g. `{"loss_mask": "bool"}`) |
| `position_ids_mode` | str | `"none"` | How to compute position_ids: `"none"`, `"doc_reset"`, `"continuous"` | | `position_ids_mode` | str | `"doc_reset"` | How to compute position_ids: `"none"`, `"doc_reset"`, `"continuous"` |
--- ---
@@ -274,12 +276,11 @@ When `sources` is set, `sections` is ignored.
### Template mode (`template: true`) ### Template mode (`template: true`)
For each message in the field's array:
1. Prepend BOS token (masked) 1. Prepend BOS token (masked)
2. Render through `chat_template` for that single message 2. For each message in the field's array:
3. Encode rendered text 1. Render through `chat_template` for that single message
4. Apply mask rule for the message's role 2. Encode rendered text
3. Apply mask rule for the message's role
### Non-template mode ### Non-template mode
@@ -287,7 +288,7 @@ Encode the field value as text. Mask value is 1 (train) or 0 (mask) per the sect
### Text config detection ### Text config detection
When no section uses `template` and all sections have `action: "train"`, the builder skips mask generation entirely — all tokens are trained. When no section uses `template` and all sections have `action: "train"`, the builder omits `loss_mask` from the output — all tokens are trained.
--- ---
@@ -298,10 +299,12 @@ When no section uses `template` and all sections have `action: "train"`, the bui
``` ```
output/ output/
__default__/ __default__/
shard_0000/
meta.json meta.json
sequence.bin sequence.bin
loss_mask.bin loss_mask.bin
wiki/ wiki/
shard_0000/
meta.json meta.json
sequence.bin sequence.bin
loss_mask.bin loss_mask.bin
@@ -324,7 +327,7 @@ output/
loss_mask.bin loss_mask.bin
``` ```
`MmapStore` discovers all shards under the domain directory via `rglob("meta.json")`. For `bin` format, `MmapStore` discovers all shards under the domain directory via `rglob("meta.json")`. For `h5` format, `H5Store` discovers `.h5`/`.hdf5` files via recursive glob.
--- ---
@@ -349,7 +352,7 @@ python scripts/tools/preprocess.py data/grpo/*.jsonl -o output/grpo/ -c configs/
from astrai.preprocessing.pipeline import Pipeline from astrai.preprocessing.pipeline import Pipeline
from astrai.config.preprocess_config import PipelineConfig from astrai.config.preprocess_config import PipelineConfig
config = PipelineConfig.from_json("sft.json") config = PipelineConfig.from_file("sft.json")
Pipeline( Pipeline(
config, config,
["data_part1.jsonl", "data_part2.jsonl"], ["data_part1.jsonl", "data_part2.jsonl"],
@@ -358,4 +361,4 @@ Pipeline(
).run() ).run()
``` ```
> Document Update Time: 2026-06-03 > Document Update Time: 2026-07-09
+28 -18
View File
@@ -58,7 +58,9 @@ on_train_begin
context.loss = loss.item() context.loss = loss.item()
stand_loss = loss / executor.grad_accum_steps stand_loss = loss / executor.grad_accum_steps
executor.backward(stand_loss) executor.backward(stand_loss)
context.iteration += 1 context.consumed_samples += (
context.config.batch_per_device * context.world_size
)
on_batch_end on_batch_end
if executor.sync_gradients: if executor.sync_gradients:
@@ -78,13 +80,13 @@ on_train_end
| `on_train_begin` | Before training starts | `GradientCheckpointingCallback` | | `on_train_begin` | Before training starts | `GradientCheckpointingCallback` |
| `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` | | `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` |
| `on_batch_begin` | Every batch | — | | `on_batch_begin` | Every batch | — |
| `on_optimizer_step` | Every accumulation window | `GradientClippingCallback`, `ValidationCallback` | | `on_optimizer_step` | Every accumulation window | `GradientClippingCallback`, `MetricCallback`, `ProgressBarCallback` |
| `on_batch_end` | Every batch | `CheckpointCallback`, `MetricLoggerCallback`, `ProgressBarCallback` | | `on_batch_end` | Every batch | `CheckpointCallback` |
| `on_epoch_end` | End of each epoch | `ProgressBarCallback` | | `on_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` |
| `on_error` | On exception during training | `CheckpointCallback`, `MetricLoggerCallback` | | `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` |
| `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricLoggerCallback`, `GradientCheckpointingCallback` | | `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricCallback`, `GradientCheckpointingCallback` |
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric_logger` (JSONL, rank-0), `progress_bar` (tqdm), `gradient_clipping`, `validation` (periodic validation on val_dataset). Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm), `gradient_clipping` (always registered; computes grad norm, clips only when `max_grad_norm` is not `None`).
## Strategies ## Strategies
@@ -106,7 +108,7 @@ $$
L_{\text{SFT}} = -\sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta) L_{\text{SFT}} = -\sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta)
$$ $$
Keys: `input_ids`, `target_ids`, `loss_mask`. Optional: `label_smoothing`. Keys: `input_ids`, `target_ids`, `loss_mask`, `position_ids`. Optional: `label_smoothing`.
### DPO (Direct Preference Optimization) ### DPO (Direct Preference Optimization)
@@ -116,21 +118,31 @@ $$
L_{\text{DPO}} = -\mathbb{E}\left[\log\sigma\left(\beta\log\frac{\pi_\theta(y_w\mid x)}{\pi_{\text{ref}}(y_w\mid x)} - \beta\log\frac{\pi_\theta(y_l\mid x)}{\pi_{\text{ref}}(y_l\mid x)}\right)\right] L_{\text{DPO}} = -\mathbb{E}\left[\log\sigma\left(\beta\log\frac{\pi_\theta(y_w\mid x)}{\pi_{\text{ref}}(y_w\mid x)} - \beta\log\frac{\pi_\theta(y_l\mid x)}{\pi_{\text{ref}}(y_l\mid x)}\right)\right]
$$ $$
Parameters: `beta=0.1`, `reduction="mean"`. Keys: `chosen`, `rejected`, `chosen_mask`, `rejected_mask`. Parameters: `beta=0.1`, `reduction="sum"`. Keys: `chosen`, `rejected`, `chosen_mask`, `rejected_mask`.
### GRPO (Group Relative Policy Optimization) ### GRPO (Group Relative Policy Optimization)
On-policy PPO with group-normalized advantages: Token-level PPO with group-normalized advantages. Advantages are derived from
scalar per-response rewards, group-normalized, and broadcast across all response
tokens. Only response tokens contribute to the loss (prompt tokens are masked
out):
$$ $$
\text{Advantage}_i = \frac{r_i - \mu}{\sigma + \epsilon} \text{Advantage}_i = \frac{r_i - \mu}{\sigma + \epsilon}
$$ $$
$$ $$
L_{\text{GRPO}} = -\mathbb{E}\left[\min\left(\frac{\pi_\theta}{\pi_{\text{ref}}}A,\; \text{clip}\left(\frac{\pi_\theta}{\pi_{\text{ref}}}, 1-\epsilon, 1+\epsilon\right)A\right)\right] + \lambda \cdot \mathbb{E}\left[(\log\pi_\theta - \log\pi_{\text{ref}})^2\right] L_{\text{GRPO}} = -\mathbb{E}_t\left[\min\left(\rho_t A,\; \text{clip}\left(\rho_t, 1-\epsilon, 1+\epsilon\right)A\right)\right] + \lambda \cdot \mathbb{E}_t\left[\frac{\pi_{\text{ref}}}{\pi_\theta} - \log\frac{\pi_{\text{ref}}}{\pi_\theta} - 1\right]
$$ $$
Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`, `sync_interval=200`, `reduction="mean"`. where $\rho_t = \pi_\theta(a_t|s_t) / \pi_{\text{old}}(a_t|s_t)$ is the
per-token importance sampling ratio against the behaviour policy
(`old_model`, synced externally between data-generation rounds) and the
expectations are over valid response tokens. The KL term regularises
$\pi_\theta$ towards a frozen reference model (`ref_model`, typically
the SFT checkpoint).
Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`. External sync of `old_model` weights via `sync_old_model()` between data-generation rounds.
Keys: `prompts`, `responses`, `masks`, `rewards`. Keys: `prompts`, `responses`, `masks`, `rewards`.
@@ -158,8 +170,8 @@ Callback wraps each `DecoderBlock.forward` with `torch.utils.checkpoint.checkpoi
## Checkpoint ## Checkpoint
``` ```
Checkpoint(state_dict, epoch, iteration, extra, meta, config) Checkpoint(state_dict, epoch, consumed_samples, extra, meta, config)
├── save(save_dir) rank-0 only: meta.json (epoch/iteration/timestamp) + config.json (model config) + model.safetensors + optional {key}.pt (optimizer.pt, scheduler.pt) ├── save(save_dir) rank-0 only: meta.json (epoch/consumed_samples/timestamp) + config.json (model config) + model.safetensors + optional {key}.pt (optimizer.pt, scheduler.pt)
└── load(save_dir, broadcast=False) loads from local disk; set broadcast=True to broadcast metadata from rank-0 └── load(save_dir, broadcast=False) loads from local disk; set broadcast=True to broadcast metadata from rank-0
``` ```
@@ -199,9 +211,7 @@ nohup python scripts/tools/train.py \
--warmup_ratio=0.05 \ --warmup_ratio=0.05 \
--max_lr=1e-4 \ --max_lr=1e-4 \
--max_grad_norm=1.0 \ --max_grad_norm=1.0 \
--adamw_beta1=0.9 \ --weight_decay=0.1 \
--adamw_beta2=0.95 \
--adamw_weight_decay=0.01 \
--window_size=2048 \ --window_size=2048 \
--ckpt_interval=10000 \ --ckpt_interval=10000 \
--ckpt_dir=./checkpoint \ --ckpt_dir=./checkpoint \
@@ -212,4 +222,4 @@ nohup python scripts/tools/train.py \
Full parameter reference at [params.md](params.md). Full parameter reference at [params.md](params.md).
> Document Update Time: 2026-05-30 > Document Update Time: 2026-07-19
+3 -5
View File
@@ -1,4 +1,4 @@
__version__ = "1.3.7" __version__ = "1.3.10"
__author__ = "ViperEkura" __author__ = "ViperEkura"
from astrai.config import ( from astrai.config import (
@@ -12,7 +12,7 @@ from astrai.config import (
from astrai.dataset import ( from astrai.dataset import (
BaseDataset, BaseDataset,
DatasetFactory, DatasetFactory,
ResumableDistributedSampler, RDSampler,
Store, Store,
StoreFactory, StoreFactory,
) )
@@ -47,7 +47,6 @@ from astrai.trainer import (
BaseScheduler, BaseScheduler,
BaseStrategy, BaseStrategy,
CallbackFactory, CallbackFactory,
Muon,
SchedulerFactory, SchedulerFactory,
StrategyFactory, StrategyFactory,
TrainCallback, TrainCallback,
@@ -75,11 +74,10 @@ __all__ = [
"GenerationRequest", "GenerationRequest",
"InferenceEngine", "InferenceEngine",
"LoRAConfig", "LoRAConfig",
"Muon",
"Pipeline", "Pipeline",
"PipelineConfig", "PipelineConfig",
"ProtocolHandler", "ProtocolHandler",
"ResumableDistributedSampler", "RDSampler",
"SamplingPipeline", "SamplingPipeline",
"SchedulerFactory", "SchedulerFactory",
"Store", "Store",
+3
View File
@@ -20,6 +20,7 @@ class BaseModelConfig(BaseConfig):
"""Base config with ``model_type`` dispatch and file I/O.""" """Base config with ``model_type`` dispatch and file I/O."""
model_type: Optional[str] = None model_type: Optional[str] = None
neftune_alpha: float = 0.0
@dataclass @dataclass
@@ -70,10 +71,12 @@ class EncoderConfig(BaseModelConfig):
rope_theta: Optional[float] = None rope_theta: Optional[float] = None
rope_scaling: Optional[dict] = None rope_scaling: Optional[dict] = None
attn_type: str = "gqa"
n_heads: Optional[int] = None n_heads: Optional[int] = None
n_kv_heads: Optional[int] = None n_kv_heads: Optional[int] = None
use_qk_norm: Optional[bool] = None use_qk_norm: Optional[bool] = None
use_gated_attention: Optional[bool] = None use_gated_attention: Optional[bool] = None
ffn_type: str = "mlp"
pooling_type: Optional[str] = None pooling_type: Optional[str] = None
normalize_embeddings: Optional[bool] = None normalize_embeddings: Optional[bool] = None
+1 -1
View File
@@ -96,7 +96,7 @@ class OutputConfig(BaseConfig):
storage_format: str = "bin" storage_format: str = "bin"
max_tokens_per_shard: int = 100_000_000 max_tokens_per_shard: int = 100_000_000
dtype: Dict[str, str] = field(default_factory=dict) dtype: Dict[str, str] = field(default_factory=dict)
position_ids_mode: str = "none" position_ids_mode: str = "doc_reset"
@dataclass @dataclass
+15 -10
View File
@@ -37,8 +37,9 @@ class TrainConfig(BaseConfig):
grad_accum_steps: int = field( grad_accum_steps: int = field(
default=1, metadata={"help": "Number of iterations between steps."} default=1, metadata={"help": "Number of iterations between steps."}
) )
max_grad_norm: float = field( max_grad_norm: Optional[float] = field(
default=1.0, metadata={"help": "Maximum gradient norm."} default=None,
metadata={"help": "Maximum gradient norm. None disables clipping."},
) )
gradient_checkpointing_modules: List[str] = field( gradient_checkpointing_modules: List[str] = field(
default_factory=list, default_factory=list,
@@ -47,14 +48,18 @@ class TrainConfig(BaseConfig):
# checkpoint setting # checkpoint setting
start_epoch: int = field(default=0, metadata={"help": "Start epoch for training."}) start_epoch: int = field(default=0, metadata={"help": "Start epoch for training."})
start_batch: int = field( start_samples: int = field(
default=0, metadata={"help": "Start batch iteration for training."} default=0,
metadata={
"help": "Start samples count (per rank). Superseded by checkpoint consumed_samples."
},
) )
ckpt_dir: str = field( ckpt_dir: str = field(
default="./checkpoint", metadata={"help": "Checkpoint directory."} default="./checkpoint", metadata={"help": "Checkpoint directory."}
) )
ckpt_interval: int = field( ckpt_interval: int = field(
default=5000, metadata={"help": "Number of iterations between checkpoints."} default=5000,
metadata={"help": "Number of optimizer steps between checkpoints."},
) )
# lora setting # lora setting
@@ -67,12 +72,8 @@ class TrainConfig(BaseConfig):
log_dir: str = field( log_dir: str = field(
default="./checkpoint/logs", metadata={"help": "Directory for metric logs."} default="./checkpoint/logs", metadata={"help": "Directory for metric logs."}
) )
log_interval: int = field(
default=100,
metadata={"help": "Number of batch iterations between metric logs."},
)
metrics: List[str] = field( metrics: List[str] = field(
default_factory=lambda: ["loss", "lr"], default_factory=lambda: ["loss", "lr", "grad_norm"],
metadata={"help": "Metrics to record during training."}, metadata={"help": "Metrics to record during training."},
) )
@@ -87,6 +88,10 @@ class TrainConfig(BaseConfig):
pin_memory: bool = field( pin_memory: bool = field(
default=False, metadata={"help": "Pin memory for dataloader."} default=False, metadata={"help": "Pin memory for dataloader."}
) )
collate_fn: Optional[Callable[[List[Any]], Any]] = field(
default=None,
metadata={"help": "Collate function for dataloader (e.g. dpo_collate_fn)."},
)
# distributed training # distributed training
nprocs: int = field( nprocs: int = field(
+14 -2
View File
@@ -1,14 +1,21 @@
from astrai.dataset.dataset import ( from astrai.dataset.dataset import (
BaseDataset, BaseDataset,
DatasetFactory, DatasetFactory,
dpo_collate_fn,
grpo_collate_fn,
) )
from astrai.dataset.sampler import ResumableDistributedSampler from astrai.dataset.sampler import RDSampler
from astrai.dataset.storage import ( from astrai.dataset.storage import (
H5Store, H5Store,
JsonlStore,
MmapStore, MmapStore,
Recordable,
Store, Store,
StoreFactory, StoreFactory,
Streamable,
detect_format, detect_format,
)
from astrai.serialization import (
load_bin, load_bin,
load_h5, load_h5,
save_bin, save_bin,
@@ -18,14 +25,19 @@ from astrai.dataset.storage import (
__all__ = [ __all__ = [
"BaseDataset", "BaseDataset",
"DatasetFactory", "DatasetFactory",
"dpo_collate_fn",
"grpo_collate_fn",
"Store", "Store",
"Streamable",
"Recordable",
"StoreFactory", "StoreFactory",
"H5Store", "H5Store",
"MmapStore", "MmapStore",
"JsonlStore",
"detect_format", "detect_format",
"save_h5", "save_h5",
"load_h5", "load_h5",
"save_bin", "save_bin",
"load_bin", "load_bin",
"ResumableDistributedSampler", "RDSampler",
] ]
+401 -180
View File
@@ -1,7 +1,31 @@
"""Dataset implementations with factory pattern for training.""" """Dataset implementations for training.
Composition over inheritance — every dataset is a thin wrapper that
binds a :class:`Store` to a particular train-type's key mapping. All
sample-id → token/record indexing lives on the Store; datasets never
know about window/stride math or segment layouts.
Class hierarchy:
BaseDataset (ABC) — holds a Store, exposes __len__/keys,
overrides __getitem__
├── SEQDataset — next-token prediction (stream)
├── SFTDataset — loss-mask + position_ids (stream)
├── DPODataset — chosen/rejected pairs (record)
└── GRPODataset — prompt + response group (record)
``DatasetFactory.load(train_type, load_path, window_size, stride, …)``
builds the Store (auto-detecting format) before constructing the
matching dataset. Passing ``store=`` skips Store construction.
When a record dataset (DPO) reads from raw JSONL, a *processor*
function (pure ``record -> Dict[str, Tensor]``) is forwarded to
:class:`JsonlStore` so tokenisation happens on the fly.
"""
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Dict, List, Optional from functools import partial
from typing import Callable, Dict, List, Optional
import torch import torch
from torch import Tensor from torch import Tensor
@@ -13,198 +37,389 @@ from astrai.dataset.storage import (
detect_format, detect_format,
) )
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.tokenize import AutoTokenizer
def dpo_tokenize(
record: dict,
tokenizer,
max_len: int = 2048,
) -> Optional[dict]:
"""Tokenize one DPO record into chosen/rejected + masks.
Applies the tokenizer's chat template so token sequences match the
SFT checkpoint's format. Prompt is rendered with
``add_generation_prompt=True``; chosen/rejected are appended as a
single assistant turn.
Accepts:
- Flat: ``{"prompt": str, "chosen": str, "rejected": str}``
- Conv: ``{"prompt": [{role, content}, ...], "chosen": [...], ...}``
- Legacy: ``{"input": str, "chosen": str, "rejected": str}``
No packing, no ``position_ids`` — DPO sequences are independent.
"""
prompt = record.get("prompt") or record.get("input")
chosen = record.get("chosen")
rejected = record.get("rejected")
if prompt is None or chosen is None or rejected is None:
return None
prompt_messages = _to_messages(prompt)
chosen_text = _extract_text(chosen)
rejected_text = _extract_text(rejected)
if chosen_text is None or rejected_text is None:
return None
chosen_messages = prompt_messages + [{"role": "assistant", "content": chosen_text}]
rejected_messages = prompt_messages + [
{"role": "assistant", "content": rejected_text}
]
prompt_ids = tokenizer.apply_chat_template(
prompt_messages, tokenize=True, add_generation_prompt=True
)
ch_ids = tokenizer.apply_chat_template(
chosen_messages, tokenize=True, add_generation_prompt=False
)
re_ids = tokenizer.apply_chat_template(
rejected_messages, tokenize=True, add_generation_prompt=False
)
full_ch = ch_ids[:max_len]
full_re = re_ids[:max_len]
prompt_len = min(len(prompt_ids), max_len)
ch_mask = [0] * prompt_len + [1] * max(0, len(full_ch) - prompt_len)
ch_mask = ch_mask[:max_len]
re_mask = [0] * prompt_len + [1] * max(0, len(full_re) - prompt_len)
re_mask = re_mask[:max_len]
return {
"chosen": full_ch,
"rejected": full_re,
"chosen_mask": ch_mask,
"rejected_mask": re_mask,
}
def _to_messages(value) -> list:
"""Accept str or conversation list; return message list."""
if isinstance(value, str):
return [{"role": "user", "content": value}]
if isinstance(value, list):
return value
return [{"role": "user", "content": str(value)}]
def _extract_text(value) -> Optional[str]:
"""Accept str or conversation list; return plain text."""
if value is None:
return None
if isinstance(value, str):
return value
if isinstance(value, list):
return "".join(m.get("content", "") for m in value if isinstance(m, dict))
return None
def dpo_processor(
record: dict,
tokenizer,
max_len: int = 2048,
) -> Dict[str, Tensor]:
"""DPO processor: wraps :func:`dpo_tokenize` and returns tensors."""
result = dpo_tokenize(record, tokenizer, max_len=max_len)
if result is None:
raise ValueError(f"Malformed DPO record: {list(record.keys())}")
return {
"chosen": torch.tensor(result["chosen"], dtype=torch.int32),
"rejected": torch.tensor(result["rejected"], dtype=torch.int32),
"chosen_mask": torch.tensor(result["chosen_mask"], dtype=torch.bool),
"rejected_mask": torch.tensor(result["rejected_mask"], dtype=torch.bool),
}
def dpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
"""Collate variable-length DPO samples into padded 2-D tensors.
Input: list of dicts, each with:
- chosen: [C_i]
- rejected: [R_i]
- chosen_mask: [C_i]
- rejected_mask: [R_i]
Output (padded to the max length across chosen/rejected within the batch):
- chosen: [B, S_max]
- rejected: [B, S_max]
- chosen_mask: [B, S_max]
- rejected_mask: [B, S_max]
"""
B = len(batch)
S_max = max(b["chosen"].size(0) for b in batch)
S_max = max(S_max, max(b["rejected"].size(0) for b in batch))
chosen = torch.zeros(B, S_max, dtype=torch.long)
rejected = torch.zeros(B, S_max, dtype=torch.long)
chosen_mask = torch.zeros(B, S_max, dtype=torch.bool)
rejected_mask = torch.zeros(B, S_max, dtype=torch.bool)
for i, b in enumerate(batch):
c_len = b["chosen"].size(0)
r_len = b["rejected"].size(0)
chosen[i, :c_len] = b["chosen"]
rejected[i, :r_len] = b["rejected"]
chosen_mask[i, :c_len] = b["chosen_mask"]
rejected_mask[i, :r_len] = b["rejected_mask"]
return {
"chosen": chosen,
"rejected": rejected,
"chosen_mask": chosen_mask,
"rejected_mask": rejected_mask,
}
def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
"""Collate variable-length GRPO samples into padded 3-D tensors.
Input: list of dicts, each with:
- prompts: [P_i]
- responses: list of G tensors, each [R_ij]
- masks: list of G tensors, each [R_ij]
- rewards: [G]
Output:
- prompts: [B, P_max]
- responses: [B, G, R_max]
- masks: [B, G, R_max]
- rewards: [B, G]
"""
B = len(batch)
G = len(batch[0]["responses"])
P_max = max(b["prompts"].size(0) for b in batch)
R_max = max(r.size(0) for b in batch for r in b["responses"])
prompts = torch.zeros(B, P_max, dtype=torch.long)
responses = torch.zeros(B, G, R_max, dtype=torch.long)
masks = torch.zeros(B, G, R_max, dtype=torch.bool)
rewards = torch.zeros(B, G, dtype=torch.float32)
for i, b in enumerate(batch):
p_len = b["prompts"].size(0)
prompts[i, :p_len] = b["prompts"]
rewards[i, : b["rewards"].size(0)] = b["rewards"]
for g in range(min(G, len(b["responses"]))):
r_len = b["responses"][g].size(0)
responses[i, g, :r_len] = b["responses"][g]
if g < len(b["masks"]):
masks[i, g, :r_len] = b["masks"][g]
return {
"prompts": prompts,
"responses": responses,
"masks": masks,
"rewards": rewards,
}
def validate_keys(store: Store, required: List[str]) -> None:
"""Raise ``KeyError`` if *store* is missing any *required* key."""
if not required:
return
actual = set(store.keys)
missing = [k for k in required if k not in actual]
if missing:
raise KeyError(
f"Store at {getattr(store, '_load_path', '?')} is missing required "
f"keys {missing}; available keys are {sorted(actual)}."
)
class BaseDataset(Dataset, ABC): class BaseDataset(Dataset, ABC):
"""Abstract base class for all dataset types. """Abstract base class for dataset types.
Implements common functionality for window-based data fetching. Holds a :class:`Store`. All sample-id indexing is delegated to the
Uses a storage abstraction for format-agnostic data loading. store — this class exposes ``__len__`` as ``len(store)`` and the
``keys`` property as ``store.keys``. Subclasses implement
``__getitem__`` with the train-type-specific key mapping and any
training-only index arithmetic (e.g. the next-token ``+1`` shift).
""" """
def __init__(self, window_size: int, stride: int): required_keys: List[str] = []
def __init__(self, store: Store):
super().__init__() super().__init__()
self.window_size = window_size self.store: Store = store
self.stride = stride validate_keys(store, self.required_keys)
self.storage: Optional[Store] = None
@property def __len__(self) -> int:
def required_keys(self) -> List[str]: return len(self.store)
"""Return required storage keys for this dataset type.
Subclasses should override to specify expected keys.
"""
return []
def _validate_keys(self):
if not self.required_keys:
return
actual_keys = set(self.storage.keys)
missing = [k for k in self.required_keys if k not in actual_keys]
if missing:
raise KeyError(
f"Dataset {type(self).__name__} requires keys {self.required_keys}, "
f"but storage at {self._load_path} only has {sorted(actual_keys)}. "
f"Missing: {missing}"
)
def load(self, load_path: str, storage_type: Optional[str] = None):
"""Load dataset from the given path.
Auto-detects the storage format if not specified.
Args:
load_path: Path to the data directory or file
storage_type: Force a specific storage type ("h5", "bin"),
or None for auto-detection
Raises:
KeyError: If the loaded storage is missing required keys.
"""
if storage_type is None:
storage_type = detect_format(load_path)
self.storage = StoreFactory.create(storage_type)
self._load_path = load_path
self.storage.load(load_path)
self._validate_keys()
@property
def count(self) -> int:
"""Return the total number of raw elements (tokens) in the dataset."""
if self.storage is None:
return 0
return len(self.storage)
@property @property
def keys(self) -> List[str]: def keys(self) -> List[str]:
"""Return the available data keys.""" return self.store.keys
if self.storage is None:
return []
return self.storage.keys
def get_index(self, index: int) -> tuple: @property
"""Calculate begin and end indices for a sample. def token_count(self) -> int:
return self.store.token_count
Args:
index: Sample index
Returns:
Tuple of (begin_idx, end_idx)
"""
if self.storage is None:
raise RuntimeError("Dataset not loaded, call load() first")
total = len(self.storage)
if total <= self.window_size:
raise ValueError(
f"Data too short: {total} tokens <= window_size {self.window_size}"
)
begin_idx = min(index * self.stride, total - 1 - self.window_size)
end_idx = min(begin_idx + self.window_size, total - 1)
return begin_idx, end_idx
@abstractmethod @abstractmethod
def __getitem__(self, index: int) -> Dict[str, Tensor]: def __getitem__(self, index: int) -> Dict[str, Tensor]:
"""Get a single sample by index.
Must be implemented by subclasses.
"""
raise NotImplementedError raise NotImplementedError
def __len__(self) -> int:
if self.storage is None:
return 0
total = len(self.storage)
if total <= self.window_size:
return 0
return (total - 1 - self.window_size) // self.stride + 1
class DatasetFactory(BaseFactory["BaseDataset"]): class DatasetFactory(BaseFactory["BaseDataset"]):
"""Factory class for creating dataset instances. """Factory for creating dataset instances by train-type.
Supports decorator-based registration for extensible dataset types. Use :meth:`DatasetFactory.register("custom")` to register new
All default dataset types (seq, sft, dpo, grpo) are registered automatically dataset classes; they must inherit from :class:`BaseDataset`.
when their classes are defined with the decorator.
Example usage:
@DatasetFactory.register("custom")
class CustomDataset(BaseDataset):
...
dataset = DatasetFactory.create("custom", window_size, stride)
""" """
@classmethod @classmethod
def load( def load(
cls, cls,
train_type: str, train_type: str,
load_path: str, load_path: Optional[str] = None,
window_size: int, window_size: int = 0,
stride: Optional[int] = None, stride: Optional[int] = None,
storage_type: Optional[str] = None, storage_type: Optional[str] = None,
tokenizer_path: Optional[str] = None,
max_len: int = 2048,
store: Optional[Store] = None,
**kwargs,
) -> "BaseDataset": ) -> "BaseDataset":
"""Create and load a dataset in one step. """Create and load a dataset in one step.
Two entry points:
- **store given**: bind it directly — the caller fully controls
Store construction and processor setup. *load_path*,
*storage_type*, *tokenizer_path*, *window_size*, *stride* are
ignored.
- **store is None**: build a Store from *load_path*, auto-detecting
format and constructing a processor when *tokenizer_path* is
given for a record dataset on JSONL.
Args: Args:
train_type: Type of training dataset train_type: Registered dataset name ("seq", "sft", "dpo",
load_path: Path to the data file "grpo", …).
window_size: Window size for data sampling load_path: Path to the data file or directory (ignored if
stride: Stride between consecutive samples (default: same as window_size) *store* is given).
storage_type: Storage type ("h5", "bin") or None for auto-detection window_size: Stream window length — only meaningful for
stream datasets (SEQ/SFT). Record datasets ignore it.
stride: Stride between consecutive stream samples
(default: same as *window_size*).
storage_type: Storage backend ("h5", "bin", "jsonl") or
None for auto-detection.
tokenizer_path: Path to tokenizer for lazy JSONL
tokenisation (record datasets only).
max_len: Max sequence length forwarded to processors.
store: Pre-built, already-loaded Store instance.
**kwargs: Extra arguments forwarded to ``store.load()``.
Returns: Returns:
Loaded dataset instance Loaded dataset instance.
""" """
if store is not None:
return cls.create(train_type, store=store)
if load_path is None:
raise ValueError("Either load_path or store must be provided")
if storage_type is None:
storage_type = detect_format(load_path)
if stride is None: if stride is None:
stride = window_size stride = window_size
dataset = cls.create(train_type, window_size, stride) processor = cls._maybe_build_processor(
dataset.load(load_path, storage_type=storage_type) train_type, storage_type, tokenizer_path, max_len
)
return dataset store_window = cls._store_window_for(train_type, window_size)
store = StoreFactory.create(
storage_type,
window_size=store_window,
stride=stride if stride else store_window,
)
if processor is not None:
store.load(load_path, processor=processor, **kwargs)
else:
store.load(load_path, **kwargs)
return cls.create(train_type, store=store)
@staticmethod
def _store_window_for(train_type: str, window_size: int) -> int:
"""Stream datasets consume ``window_size``; record datasets ignore it.
Record datasets (dpo/grpo) treat each record as an independent
training unit and never window, so the store is built with
``window_size=0`` and ``len(store)`` returns the record count.
"""
if train_type in ("seq", "sft"):
return window_size
return 0
@staticmethod
def _maybe_build_processor(
train_type: str,
storage_type: str,
tokenizer_path: Optional[str],
max_len: int,
) -> Optional[Callable[[dict], Dict[str, Tensor]]]:
"""Build an on-the-fly tokenisation processor if applicable.
Only raw JSONL + record datasets (DPO/GRPO) need a processor;
pre-tokenised backends (H5/bin) and stream datasets (SEQ/SFT)
return ``None`` so no tokenizer is loaded.
"""
if tokenizer_path is None or storage_type != "jsonl":
return None
if train_type == "dpo":
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
return None
@DatasetFactory.register("seq") @DatasetFactory.register("seq")
class SEQDataset(BaseDataset): class SEQDataset(BaseDataset):
"""Dataset for sequential next-token prediction training.""" """Dataset for sequential next-token prediction training.
@property Stream mode: ``store.fetch(begin, end, "sequence")`` returns the
def required_keys(self) -> List[str]: input window; the +1 shifted call returns the next-token target.
return ["sequence"] """
def _fetch_data(self, begin_idx: int, end_idx: int) -> Tensor: required_keys = ["sequence"]
return self.storage.fetch(begin_idx, end_idx, "sequence")
def __getitem__(self, index): def __getitem__(self, index: int):
begin_idx, end_idx = self.get_index(index) begin, end = self.store.sample_window(index)
x = self.store.fetch(begin, end, "sequence")
x = self._fetch_data(begin_idx, end_idx).to(dtype=torch.long) y = self.store.fetch(begin + 1, end + 1, "sequence")
y = self._fetch_data(begin_idx + 1, end_idx + 1).to(dtype=torch.long) return {
"input_ids": x.to(dtype=torch.long),
return {"input_ids": x, "target_ids": y} "target_ids": y.to(dtype=torch.long),
}
@DatasetFactory.register("sft") @DatasetFactory.register("sft")
class SFTDataset(BaseDataset): class SFTDataset(BaseDataset):
"""Dataset for supervised fine-tuning with loss masking.""" """Dataset for supervised fine-tuning with loss masking.
@property Stream mode: ``sequence``/``loss_mask``/``position_ids`` are sliced
def required_keys(self) -> List[str]: to the window. ``loss_mask`` and ``target_ids`` use the +1 shifted
return ["sequence", "loss_mask", "position_ids"] slice so they align with the predicted positions.
"""
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor: required_keys = ["sequence", "loss_mask", "position_ids"]
return self.storage.fetch(begin_idx, end_idx, key)
def __getitem__(self, index):
begin_idx, end_idx = self.get_index(index)
x = self._fetch_data(begin_idx, end_idx, "sequence")
y = self._fetch_data(begin_idx + 1, end_idx + 1, "sequence")
position_ids = self._fetch_data(begin_idx, end_idx, "position_ids")
loss_mask = self._fetch_data(begin_idx + 1, end_idx + 1, "loss_mask")
def __getitem__(self, index: int):
begin, end = self.store.sample_window(index)
x = self.store.fetch(begin, end, "sequence")
y = self.store.fetch(begin + 1, end + 1, "sequence")
position_ids = self.store.fetch(begin, end, "position_ids")
loss_mask = self.store.fetch(begin + 1, end + 1, "loss_mask")
return { return {
"input_ids": x.to(dtype=torch.long), "input_ids": x.to(dtype=torch.long),
"target_ids": y.to(dtype=torch.long), "target_ids": y.to(dtype=torch.long),
@@ -215,59 +430,65 @@ class SFTDataset(BaseDataset):
@DatasetFactory.register("dpo") @DatasetFactory.register("dpo")
class DPODataset(BaseDataset): class DPODataset(BaseDataset):
"""Dataset for Direct Preference Optimization training.""" """Record-structured dataset for Direct Preference Optimization.
@property Each sample is one preference pair (chosen + rejected) and is an
def required_keys(self) -> List[str]: independent training unit — no windowing, stride, or cross-record
return ["chosen", "rejected", "chosen_mask", "rejected_mask"] concatenation. This keeps each sequence self-contained so attention
never leaks across preference pairs.
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor: Two loading paths (handled by :class:`DatasetFactory`):
return self.storage.fetch(begin_idx, end_idx, key)
def __getitem__(self, index: int): - **Pre-tokenized** (H5/bin): ``store.load(path)`` reads per-record
begin_idx, end_idx = self.get_index(index) tensors; ``__getitem__`` returns them directly.
- **Raw JSONL** (``tokenizer_path=...``): builds a lazy processor
via :func:`dpo_processor` that tokenises on the fly — no packing,
no ``position_ids``.
"""
chosen = self._fetch_data(begin_idx, end_idx, "chosen").to(dtype=torch.long) required_keys = ["chosen", "rejected", "chosen_mask", "rejected_mask"]
rejected = self._fetch_data(begin_idx, end_idx, "rejected").to(dtype=torch.long)
chosen_mask = self._fetch_data(begin_idx, end_idx, "chosen_mask").to(
dtype=torch.bool
)
rejected_mask = self._fetch_data(begin_idx, end_idx, "rejected_mask").to(
dtype=torch.bool
)
def make_processor(self, tokenizer, max_len: int):
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
def __getitem__(self, index: int) -> Dict[str, Tensor]:
return { return {
"chosen": chosen, "chosen": self.store.fetch_record(index, "chosen").to(dtype=torch.long),
"rejected": rejected, "rejected": self.store.fetch_record(index, "rejected").to(dtype=torch.long),
"chosen_mask": chosen_mask, "chosen_mask": self.store.fetch_record(index, "chosen_mask").to(
"rejected_mask": rejected_mask, dtype=torch.bool
),
"rejected_mask": self.store.fetch_record(index, "rejected_mask").to(
dtype=torch.bool
),
} }
@DatasetFactory.register("grpo") @DatasetFactory.register("grpo")
class GRPODataset(BaseDataset): class GRPODataset(BaseDataset):
"""Dataset for Group Relative Policy Optimization training.""" """Dataset for offline Group Relative Policy Optimization.
@property Each sample is one prompt with its group of responses and scalar
def required_keys(self) -> List[str]: rewards — an independent training unit with no windowing or stride.
return ["prompts", "responses", "masks", "rewards"]
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor: Expected storage layout (produced by JsonlStore or pre-tokenized):
return self.storage.fetch(begin_idx, end_idx, key)
- ``prompts``: List[Tensor] — one 1-D token tensor per record
- ``responses``: List[List[Tensor]] — G response tensors per record
- ``masks``: List[List[Tensor]] — G mask tensors per record
- ``rewards``: List[Tensor] — one 1-D float tensor (len G) per record
"""
required_keys = ["prompts", "responses", "masks", "rewards"]
def __getitem__(self, index: int) -> Dict[str, Tensor]: def __getitem__(self, index: int) -> Dict[str, Tensor]:
begin_idx, end_idx = self.get_index(index) prompts = self.store.fetch_record(index, "prompts")
responses = self.store.fetch_record(index, "responses")
prompts = self._fetch_data(begin_idx, end_idx, "prompts").to(dtype=torch.long) masks = self.store.fetch_record(index, "masks")
responses = self._fetch_data(begin_idx, end_idx, "responses").to( rewards = self.store.fetch_record(index, "rewards")
dtype=torch.long
)
masks = self._fetch_data(begin_idx, end_idx, "masks").to(dtype=torch.bool)
rewards = self._fetch_data(begin_idx, end_idx, "rewards")
return { return {
"prompts": prompts, "prompts": prompts.to(dtype=torch.long),
"responses": responses, "responses": [r.to(dtype=torch.long) for r in responses],
"masks": masks, "masks": [m.to(dtype=torch.bool) for m in masks],
"rewards": rewards, "rewards": rewards.to(dtype=torch.float32),
} }
+10 -1
View File
@@ -5,7 +5,15 @@ import torch.distributed as dist
from torch.utils.data import Dataset, Sampler from torch.utils.data import Dataset, Sampler
class ResumableDistributedSampler(Sampler[int]): class RDSampler(Sampler[int]):
"""Resumable Distributed Sampler.
A distributed sampler that supports checkpoint-based resume: iteration
state (epoch, position) is tracked so training can continue from the
exact sample after a restart. Shards the dataset across
``dist.world_size`` replicas with optional shuffling.
"""
def __init__( def __init__(
self, self,
data_source: Dataset, data_source: Dataset,
@@ -74,6 +82,7 @@ class ResumableDistributedSampler(Sampler[int]):
self.epoch += 1 self.epoch += 1
self._indices = None self._indices = None
self.iter = self.iter % self.num_samples_per_replica
@property @property
def _remaining(self): def _remaining(self):
+499 -129
View File
@@ -1,98 +1,70 @@
"""Storage backends for different data formats. """Storage backends for different data formats.
Layers: Architecture (composition over inheritance):
- I/O layer: save_* / load_* functions, read/write raw files (HDF5/bin)
return Dict[str, List[Tensor]] — format-specific, no state
- Store (ABC): central abstraction, normalizes multi-segment into
Dict[str, List[Tensor]] per key via _normalize(),
fetch() uses bisect across segments — no forced concat
- Dataset layer: BaseDataset owns a Store, only calls store.fetch(begin, end, key)
Key properties: Store (ABC) — owns _data/_cum/_offsets bookkeeping
- Multi-segment: segments kept as-is, no forced concatenation — safe for + window_size/stride for sample-id
datasets larger than RAM indexing. __getitem__/__len__ produce
- Explicit length: _length = min(total elements across keys), set at load, the smallest iterable unit so Dataset
__len__ returns O(1) classes are pure delegators.
- Zero-copy mmap: MmapStore wraps np.memmap(mode="r"), all DataLoader Streamable (mixin) — raw token slice fetch(begin, end, keys)
workers share OS page-cache pages Recordable (mixin) — raw record slice fetch_record(idx, keys)
H5Store(Store, Streamable, Recordable)
MmapStore(Store, Streamable, Recordable)
JsonlStore(Store, Streamable, Recordable)
Each mixin is a stateless trait that relies on ``self._data`` etc.
provided by :class:`Store`. Concrete stores mix in whichever access
primitives they support — ``Store`` is the sole base class, so there is
no diamond inheritance or MRO ambiguity.
Sample-id indexing lives on :class:`Store`, not on the dataset:
- **Stream mode** (``window_size > 0``): ``len(store)`` returns the number
of ``(window_size, stride)`` windows that fit in the token river;
``store[i]`` returns the *i*-th window as a dict of per-key tensors;
``store.sample_window(i)`` exposes the underlying ``(begin, end)``
token slice for callers (e.g. next-token trainers) that need a +1
shifted companion window.
- **Record mode** (``num_records > 0``): ``len(store)`` returns the
record count; ``store[i]`` returns the *i*-th record dict.
Raw token/record access via :meth:`fetch` / :meth:`fetch_record`
remains available for low-level callers that want explicit index
control. ``store.token_count`` is the total stream token count (what
``len(store)`` used to mean in the legacy stream-only API).
``segments_are_records`` (class attribute on each Store subclass)
tells ``_normalize`` whether segments are inherently per-record (H5/
JSONL) or opaque shards (bin). Record access for bin relies on
``_offsets`` instead.
:class:`JsonlStore` supports a lazy mode (``processor=fn``) that keeps
raw records and defers tokenisation to ``fetch_record`` — used by DPO
to train directly from a ``.jsonl`` file without a pre-tokenised copy.
""" """
import bisect import bisect
import glob import glob
import json import json
import os import logging
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from pathlib import Path from pathlib import Path
from typing import Dict, List, Union from typing import Callable, Dict, List, Optional, Tuple, Union
import h5py
import numpy as np
import torch import torch
from torch import Tensor from torch import Tensor
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.preprocessing.transform import TokenizeTransform
from astrai.serialization import (
load_bin,
load_bin_offsets,
load_h5,
)
logger = logging.getLogger(__name__)
def save_h5(file_path: str, file_name: str, tensor_group: Dict[str, List[Tensor]]):
os.makedirs(file_path, exist_ok=True)
full_file_path = os.path.join(file_path, f"{file_name}.h5")
with h5py.File(full_file_path, "w") as f:
for key, tensors in tensor_group.items():
grp = f.create_group(key)
for idx, tensor in enumerate(tensors):
arr = tensor.cpu().numpy()
grp.create_dataset(f"data_{idx}", data=arr)
def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
tensor_group: Dict[str, List[Tensor]] = {}
root_path = Path(file_path)
h5_files = list(root_path.rglob("*.h5")) + list(root_path.rglob("*.hdf5"))
for h5_file in h5_files:
with h5py.File(h5_file, "r") as f:
for key in f.keys():
grp = f[key]
dsets = []
for dset_name in grp.keys():
dset = grp[dset_name]
tensor = torch.from_numpy(dset[:])
if share_memory:
tensor = tensor.share_memory_()
dsets.append(tensor)
if tensor_group.get(key) is None:
tensor_group[key] = []
tensor_group[key].extend(dsets)
return tensor_group
def save_bin(file_path: str, tensor_group: Dict[str, List[Tensor]]):
os.makedirs(file_path, exist_ok=True)
meta = {}
for key, tensors in tensor_group.items():
cat = torch.cat(tensors, dim=0)
meta[key] = {"shape": list(cat.shape), "dtype": str(cat.dtype).split(".")[-1]}
np.asarray(cat.cpu().numpy()).tofile(os.path.join(file_path, f"{key}.bin"))
with open(os.path.join(file_path, "meta.json"), "w") as f:
json.dump(meta, f)
def load_bin(file_path: str) -> Dict[str, List[Tensor]]:
with open(os.path.join(file_path, "meta.json"), "r") as f:
meta = json.load(f)
segments: Dict[str, List[Tensor]] = {}
for key, info in meta.items():
arr = np.memmap(
os.path.join(file_path, f"{key}.bin"),
dtype=info["dtype"],
mode="r+",
shape=tuple(info["shape"]),
)
segments[key] = [torch.from_numpy(arr)]
return segments
def detect_format(load_path: str) -> str: def detect_format(load_path: str) -> str:
@@ -102,7 +74,7 @@ def detect_format(load_path: str) -> str:
load_path: Directory or file path load_path: Directory or file path
Returns: Returns:
Format string ("h5" or "bin") Format string ("h5", "bin", "jsonl", or "processed")
Raises: Raises:
FileNotFoundError: If no supported data files are found FileNotFoundError: If no supported data files are found
@@ -112,6 +84,8 @@ def detect_format(load_path: str) -> str:
suffix = root.suffix.lower() suffix = root.suffix.lower()
if suffix in (".h5", ".hdf5"): if suffix in (".h5", ".hdf5"):
return "h5" return "h5"
if suffix == ".jsonl":
return "jsonl"
raise ValueError(f"Unsupported file format: {suffix}") raise ValueError(f"Unsupported file format: {suffix}")
h5_files = [ h5_files = [
@@ -128,54 +102,259 @@ def detect_format(load_path: str) -> str:
) > 0 ) > 0
if has_meta: if has_meta:
return "bin" return "bin"
jsonl_files = [
Path(p) for p in glob.glob(str(root / "**" / "*.jsonl"), recursive=True)
]
if jsonl_files:
return "jsonl"
raise FileNotFoundError(f"No supported data files found at {load_path}") raise FileNotFoundError(f"No supported data files found at {load_path}")
class Store(ABC): class Store(ABC):
"""String keys -> segmented tensors with ``fetch(begin, end, keys)``. """Common base for all storage backends.
Each key maps to one or more tensor segments (no forced concatenation). A Store owns both its data layout AND its sample-id → token/record
``len(store)`` returns ``self._length`` (explicit, O(1)), the minimum index translation. Datasets are thin wrappers that bind a Store
total element count across all keys. to a particular train-type's key mapping; they never know about
window/stride math.
Subclasses fill ``self._data`` and ``self._cum`` during ``load()`` Two iteration modes:
via ``_normalize()``.
- **Stream** (``window_size > 0``): data is treated as one long
token river. ``len(store)`` returns the number of windows;
``store[i]`` slices every stream-compatible key to window ``i``;
``store.sample_window(i)`` returns the ``(begin, end)`` token
slice for callers needing a +1 shifted companion window.
- **Record** (``num_records > 0``): data is per-record.
``len(store)`` returns ``num_records``; ``store[i]`` returns
the *i*-th record as a dict.
Raw token slicing is still available via :meth:`fetch` (mixed in
by :class:`Streamable`) when a store has stream support configured.
Raw record slicing via :meth:`fetch_record` (mixed in by
:class:`Recordable`) when a store has record support.
``token_count`` exposes the raw total stream length — this is what
``len(store)`` returned in the legacy stream-only API and what
stream-bound ``fetch`` uses for its bounds check.
""" """
def __init__(self): segments_are_records: bool = False
def __init__(
self,
window_size: int = 0,
stride: Optional[int] = None,
):
self._data: Dict[str, List[Tensor]] = {} self._data: Dict[str, List[Tensor]] = {}
self._cum: Dict[str, List[int]] = {} self._cum: Dict[str, List[int]] = {}
self._offsets: Dict[str, List[int]] = {}
self._length: int = 0 self._length: int = 0
self._num_records: int = 0
self._window_size: int = int(window_size)
self._stride: int = int(stride) if stride is not None else int(window_size)
@abstractmethod @abstractmethod
def load(self, path: str) -> None: def load(self, path: str, **kwargs) -> None:
raise NotImplementedError raise NotImplementedError
@property @property
def keys(self) -> List[str]: def keys(self) -> List[str]:
return list(self._data.keys()) return list(self._data.keys())
def __len__(self) -> int: @property
def window_size(self) -> int:
return self._window_size
@property
def stride(self) -> int:
return self._stride
@property
def token_count(self) -> int:
"""Total tokens across all stream segments.
Useful for the bounds-checked raw :meth:`fetch` and as the
legacy ``len(store)`` value.
"""
return self._length return self._length
@property
def num_records(self) -> int:
"""Number of records available via :meth:`fetch_record`.
Non-zero only when the backing layout provides per-record
indexing (H5/JSONL segments or bin ``_offsets``).
"""
return self._num_records
@property
def num_samples(self) -> int:
"""Number of items produced by ``__getitem__``.
Stream-mode wins when ``window_size > 0`` and there are tokens
to slice; otherwise falls back to ``num_records``.
"""
if self._window_size > 0 and self._length > 0:
total = self._length
w = self._window_size
if total <= w:
return 0
return (total - 1 - w) // self._stride + 1
return self._num_records
def __len__(self) -> int:
return self.num_samples
def __getitem__(self, index: int) -> Dict[str, Tensor]:
if index < 0:
index += self.num_samples
if not 0 <= index < self.num_samples:
raise IndexError(
f"Store index out of range: {index}, num_samples={self.num_samples}"
)
if self._window_size > 0 and self._length > 0:
begin, end = self.sample_window(index)
keys = self._stream_keys()
return {k: self.fetch(begin, end, k) for k in keys}
return self.fetch_record(index, self._record_keys())
def sample_window(self, index: int) -> Tuple[int, int]:
"""Return ``(begin, end)`` token positions for stream sample *index*.
The clipped tail keeps the last reachable window inside the
token river instead of overshooting. Caller is responsible
for staying within :attr:`num_samples`: an out-of-range index
raises ``IndexError``.
"""
if self._window_size <= 0:
raise RuntimeError("sample_window() requires window_size > 0 (stream mode)")
if self._window_size <= 0 or self._length <= self._window_size:
raise IndexError(
f"Data too short for window: token_count={self._length}, "
f"window_size={self._window_size}"
)
if not 0 <= index < self.num_samples:
raise IndexError(
f"Sample index out of range: {index}, num_samples={self.num_samples}"
)
total = self._length
begin = min(index * self._stride, total - 1 - self._window_size)
end = min(begin + self._window_size, total - 1)
return begin, end
def _stream_keys(self) -> List[str]:
out: List[str] = []
for k, tensors in self._data.items():
if tensors and isinstance(tensors[0], list):
continue
out.append(k)
return out
def _record_keys(self) -> List[str]:
return list(self._data.keys())
def _normalize(
self,
raw: Dict[str, list],
offsets: Optional[Dict[str, List[int]]] = None,
):
"""Register segments and pre-compute indices for both access modes.
Stream mode: ``_cum[key]`` accumulates per-segment lengths so
``Streamable._fetch_stream_key`` can bisect across segments
without concatenation.
Record mode: if *offsets* is provided (bin layout),
``_offsets[key]`` stores cumulative per-record offsets into the
single concatenated segment. Otherwise, when
``segments_are_records`` is True (H5/JSONL), ``_data[key]`` is
a per-record list and ``fetch_record`` indexes it directly.
Nested keys (GRPO ``responses``/``masks`` as
``List[List[Tensor]]``) are stored as-is and excluded from both
cumulative bookkeepings — they are only accessed record-by-record.
"""
flat_lengths = []
for key, tensors in raw.items():
self._data[key] = tensors
if not tensors:
self._cum[key] = []
flat_lengths.append(0)
continue
if isinstance(tensors[0], list):
self._cum[key] = []
continue
cum = []
total = 0
for t in tensors:
total += t.shape[0]
cum.append(total)
self._cum[key] = cum
flat_lengths.append(cum[-1] if cum else 0)
self._length = min(flat_lengths) if flat_lengths else 0
valid_offsets: Dict[str, List[int]] = {}
if offsets:
for key, off in offsets.items():
segs = self._data.get(key, [])
if len(segs) == 1 and len(off) > 1:
valid_offsets[key] = off
elif len(segs) > 1:
logger.warning(
"Key '%s' has %d segments with offsets — record mode "
"disabled for this key (multi-shard bin+offsets not "
"supported). Merge shards or use H5/JSONL.",
key,
len(segs),
)
self._offsets = valid_offsets
if valid_offsets:
record_counts = [len(v) - 1 for v in valid_offsets.values()]
self._num_records = min(record_counts) if record_counts else 0
elif self.segments_are_records:
per_record_counts = []
for key, tensors in self._data.items():
if tensors and isinstance(tensors[0], list):
continue
per_record_counts.append(len(tensors))
self._num_records = min(per_record_counts) if per_record_counts else 0
else:
self._num_records = 0
class Streamable:
"""Mixin granting raw token-stream access via :meth:`fetch`.
Stateless trait relying on ``self._data``, ``self._cum``,
``self._length`` maintained by :class:`Store`. Stream mode is
active when the owning store has ``window_size > 0``; for stores
that can also serve record access (H5/JSONL/bin+offsets), the
``fetch_record`` API from :class:`Recordable` is used instead.
"""
def fetch( def fetch(
self, self,
begin: int, begin: int,
end: int, end: int,
keys: Union[str, List[str]], keys: Union[str, List[str]],
): ):
if not self._data: return _stream_fetch(self, begin, end, keys)
def _stream_fetch(self, begin: int, end: int, keys: Union[str, List[str]]):
if not getattr(self, "_data", None):
raise RuntimeError("Store not loaded") raise RuntimeError("Store not loaded")
if not (0 <= begin < self._length and 0 <= end <= self._length): if not (0 <= begin < self._length and 0 <= end <= self._length):
raise ValueError( raise ValueError(
f"Index out of bounds: begin={begin}, end={end}, length={self._length}" f"Index out of bounds: begin={begin}, end={end}, length={self._length}"
) )
if isinstance(keys, str): if isinstance(keys, str):
return self._fetch_key(keys, begin, end) return _fetch_stream_key(self, keys, begin, end)
return {k: self._fetch_key(k, begin, end) for k in keys} return {k: _fetch_stream_key(self, k, begin, end) for k in keys}
def _fetch_key(self, key: str, begin: int, end: int) -> Tensor:
"""Fetch slice [begin, end) across potentially multiple segments.""" def _fetch_stream_key(self, key: str, begin: int, end: int) -> Tensor:
segments = self._data[key] segments = self._data[key]
cum = self._cum[key] cum = self._cum[key]
seg_start = bisect.bisect_right(cum, begin) seg_start = bisect.bisect_right(cum, begin)
@@ -190,77 +369,268 @@ class Store(ABC):
return results[0] if len(results) == 1 else torch.cat(results, dim=0) return results[0] if len(results) == 1 else torch.cat(results, dim=0)
def _normalize(self, raw: Dict[str, List[Tensor]]):
"""Register segments and pre-compute cumulative lengths.
Does NOT concatenate — segments are kept as-is to avoid OOM on class Recordable:
large datasets. Sets ``self._length`` to the minimum total """Mixin granting raw record access via :meth:`fetch_record`.
element count across all keys.
Stateless trait relying on ``self._data``, ``self._offsets``,
``self._num_records`` maintained by :class:`Store`.
""" """
for key, tensors in raw.items():
self._data[key] = tensors def fetch_record(
cum = [] self,
total = 0 index: int,
for t in tensors: keys: Union[str, List[str]],
total += t.shape[0] ):
cum.append(total) return _record_fetch(self, index, keys)
self._cum[key] = cum
self._length = (
min((cum[-1] if cum else 0) for cum in self._cum.values()) def _record_fetch(self, index: int, keys: Union[str, List[str]]):
if self._cum if not getattr(self, "_data", None) and self._num_records == 0:
else 0 raise RuntimeError("Store not loaded")
if not 0 <= index < self._num_records:
raise ValueError(
f"Record index out of bounds: {index}, num_records={self._num_records}"
) )
if isinstance(keys, str):
return _fetch_record_key(self, keys, index)
return {k: _fetch_record_key(self, k, index) for k in keys}
def _fetch_record_key(self, key: str, index: int):
offsets = self._offsets.get(key)
if offsets:
start = offsets[index]
end = (
offsets[index + 1]
if index + 1 < len(offsets)
else self._data[key][0].shape[0]
)
return self._data[key][0][start:end]
return self._data[key][index]
class StoreFactory(BaseFactory["Store"]): class StoreFactory(BaseFactory["Store"]):
"""Factory for creating Store instances by type name. """Factory for creating Store instances by type name."""
Example::
@StoreFactory.register("custom")
class CustomStore(Store):
...
"""
@StoreFactory.register("h5") @StoreFactory.register("h5")
class H5Store(Store): class H5Store(Store, Streamable, Recordable):
"""HDF5-based storage backend (pre-tokenized data).""" """HDF5-based storage backend (pre-tokenized data).
def load(self, path: str): Each key is stored as a group of per-record datasets (``data_0``,
``data_1``, …). Supports both access modes:
- **Stream**: ``fetch(begin, end, key)`` and ``store[i]`` slice
across concatenated records via ``_cum`` — used by SEQ/SFT.
- **Record**: ``fetch_record(i, key)`` and ``store[i]`` (when
``window_size == 0``) index ``_data[key]`` directly — used by
DPO/GRPO.
"""
segments_are_records = True
def __init__(
self,
window_size: int = 0,
stride: Optional[int] = None,
):
super().__init__(window_size=window_size, stride=stride)
def load(self, path: str, **kwargs):
self._normalize(load_h5(path)) self._normalize(load_h5(path))
@StoreFactory.register("bin") @StoreFactory.register("bin")
class MmapStore(Store): class MmapStore(Store, Streamable, Recordable):
"""Memory-mapped binary storage backend. """Memory-mapped binary storage backend.
Each key is a single .bin file backed by ``np.memmap(mode="r")``. Each key is a single .bin file backed by ``np.memmap(mode="r")``.
No per-process memory duplication — all DataLoader workers share the No per-process memory duplication — all DataLoader workers share the
same OS page-cache pages. same OS page-cache pages.
Format on disk:: Supports both access modes:
data_root/ - **Stream**: always available via :meth:`fetch`.
meta.json # {key: {shape, dtype}, ...} - **Record** (``fetch_record(i, key)``): only when ``meta.json``
<key>.bin # raw numpy array, one per key contains per-record ``offsets`` (written via
``save_bin(..., record_keys=...)``). Legacy bin files without
offsets have ``num_records == 0`` and ``len(store)`` reflects the
windowed sample count when ``window_size > 0``.
``segments_are_records`` is ``False`` here (bin segments are
contiguous streams, not per-record) — record access is driven
purely by ``_offsets``.
""" """
def load(self, path: str): segments_are_records = False
def __init__(
self,
window_size: int = 0,
stride: Optional[int] = None,
):
super().__init__(window_size=window_size, stride=stride)
self._mmap_refs: List[Tensor] = []
def load(self, path: str, **kwargs):
self._mmap_refs = [] self._mmap_refs = []
root = Path(path) root = Path(path)
all_raw: Dict[str, List[Tensor]] = {} all_raw: Dict[str, List[Tensor]] = {}
all_offsets: Dict[str, List[int]] = {}
meta_paths = [ meta_paths = [
Path(p) for p in glob.glob(str(root / "**" / "meta.json"), recursive=True) Path(p) for p in glob.glob(str(root / "**" / "meta.json"), recursive=True)
] ]
for meta_path in meta_paths: for meta_path in meta_paths:
raw = load_bin(str(meta_path.parent)) raw = load_bin(str(meta_path.parent))
off = load_bin_offsets(str(meta_path.parent))
for key, tensors in raw.items(): for key, tensors in raw.items():
if key not in all_raw: if key not in all_raw:
all_raw[key] = [] all_raw[key] = []
all_raw[key].extend(tensors) all_raw[key].extend(tensors)
for key, o in off.items():
if key not in all_offsets:
all_offsets[key] = []
all_offsets[key].extend(o)
if not meta_paths: if not meta_paths:
raise FileNotFoundError(f"No meta.json found under {path}") raise FileNotFoundError(f"No meta.json found under {path}")
self._normalize(all_raw) self._normalize(all_raw, offsets=all_offsets or None)
for tensors in self._data.values(): for tensors in self._data.values():
self._mmap_refs.extend(tensors) self._mmap_refs.extend(tensors)
class JsonlSource:
"""Read raw JSON records from a ``.jsonl`` file or directory.
A thin reader used by :class:`JsonlStore` in processor mode — holds
no tokenizer, performs no tokenisation, just yields dicts.
"""
def __init__(self, path: str):
self.path = Path(path)
self._records: Optional[List[dict]] = None
def load(self) -> List[dict]:
if self._records is None:
self._records = self._read(self.path)
return self._records
@staticmethod
def _read(root: Path) -> List[dict]:
if root.is_file():
return JsonlSource._read_file(root)
return JsonlSource._read_dir(root)
@staticmethod
def _read_file(path: Path) -> List[dict]:
records: List[dict] = []
with open(path, "r", encoding="utf-8") as f:
for line in f:
line = line.strip()
if not line:
continue
try:
records.append(json.loads(line))
except json.JSONDecodeError:
logger.warning("Failed to parse JSON line in %s, skipping", path)
return records
@staticmethod
def _read_dir(root: Path) -> List[dict]:
records: List[dict] = []
for jsonl_path in sorted(root.glob("*.jsonl")):
records.extend(JsonlSource._read_file(jsonl_path))
return records
@StoreFactory.register("jsonl")
class JsonlStore(Store, Streamable, Recordable):
"""JSONL reader with two 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:
- **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.
- **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.
"""
CONFIG_NAME = "dataset_config.json"
segments_are_records = True
def __init__(
self,
window_size: int = 0,
stride: Optional[int] = None,
):
super().__init__(window_size=window_size, stride=stride)
self._source: Optional[JsonlSource] = None
self._processor: Optional[Callable[[dict], Dict[str, Tensor]]] = None
self._keys_cache: Optional[List[str]] = None
def load(self, path: str, transform=None, processor=None, **kwargs):
self._source = JsonlSource(path)
records = self._source.load()
if processor is not None:
self._processor = processor
self._num_records = len(records)
return
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))
transformed = transform.apply(records)
self._normalize(transformed)
@property
def keys(self) -> List[str]:
if self._processor is not None:
if self._keys_cache is None and self._num_records > 0:
sample = self._processor(self._source.load()[0])
self._keys_cache = list(sample.keys())
return self._keys_cache or []
return list(self._data.keys())
def fetch_record(self, index: int, keys: Union[str, List[str]]):
if self._processor is not None:
if not 0 <= index < self._num_records:
raise ValueError(
f"Record index out of bounds: {index}, "
f"num_records={self._num_records}"
)
record = self._source.load()[index]
data = self._processor(record)
if isinstance(keys, str):
return data[keys]
return {k: data[k] for k in keys}
return _record_fetch(self, index, keys)
def fetch(self, begin: int, end: int, keys: Union[str, List[str]]):
if self._processor is not None:
raise RuntimeError(
"JsonlStore in lazy (processor) mode does not support "
"stream fetch(); use fetch_record() instead."
)
return _stream_fetch(self, begin, end, keys)
def __getitem__(self, index: int) -> Dict[str, Tensor]:
if self._processor is not None:
return self.fetch_record(index, self._record_keys())
return super().__getitem__(index)
+29
View File
@@ -0,0 +1,29 @@
"""CUDA attention kernel wrappers with torch fallback.
Public API:
- ``attn_decode`` — single-query decode attention
- ``attn_prefill`` — multi-query prefill attention
- ``attn_paged_decode`` — paged decode attention (direct page-table access)
Interface (shared by all wrappers):
causal_offset: -1 = non-causal; >=0 = absolute position of first Q token
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True = keep)
scale: 0.0 = auto (1/sqrt(head_dim)); >0 = explicit
layout: "bhld" (default) or "blhd"
Causal and mask can coexist — both are applied simultaneously.
Each wrapper dispatches to its compiled CUDA kernel (``astrai.extension.attn_*``)
when available, otherwise falls back to ``torch.nn.functional.scaled_dot_product_attention``.
"""
from astrai.extension.loader import KERNEL_NAMES, is_available
from astrai.extension.ops import attn_decode, attn_paged_decode, attn_prefill
__all__ = [
"attn_decode",
"attn_paged_decode",
"attn_prefill",
"is_available",
"KERNEL_NAMES",
]
+36
View File
@@ -0,0 +1,36 @@
"""Dynamic discovery and loading of compiled CUDA kernel modules.
Each kernel is registered in ``csrc/build.py`` and built into a ``.so`` placed
in this package directory. On import we try to load each one; kernels that
failed to build (or are running on a CPU-only machine) are marked unavailable
so the wrapper functions can fall back to ``torch`` SDPA.
"""
import importlib
import logging
logger = logging.getLogger(__name__)
KERNEL_NAMES = ["attn_decode", "attn_prefill", "attn_paged_decode"]
_available: dict[str, bool] = {}
_modules: dict[str, object] = {}
for _name in KERNEL_NAMES:
try:
_mod = importlib.import_module(f".{_name}", package=__package__)
_available[_name] = True
_modules[_name] = _mod
except ImportError:
_available[_name] = False
_modules[_name] = None
def is_available(name: str) -> bool:
"""Return ``True`` if the compiled kernel ``name`` was loaded."""
return _available.get(name, False)
def get_module(name: str) -> object:
"""Return the loaded kernel module for ``name``, or ``None`` if unavailable."""
return _modules.get(name)
+246
View File
@@ -0,0 +1,246 @@
"""GQA attention wrapper functions — one entry point per compiled kernel.
Each wrapper dispatches to its CUDA kernel (loaded in ``loader.py``) when
available, otherwise falls back to ``torch`` SDPA.
Interface (all functions):
causal_offset: -1 = non-causal; >=0 = absolute position of first Q token
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool)
scale: 0.0 = auto (1/sqrt(head_dim)); >0 = explicit
layout: "bhld" (default) or "blhd"
Add new kernel wrappers here; split into per-variant files only if this file
grows large.
"""
import math
import torch
import torch.nn.functional as F
from astrai.extension.loader import _available, _modules
_LAYOUT_CODES: dict[str, int] = {"bhld": 0, "blhd": 1}
def _parse_layout(layout: str | int) -> int:
if isinstance(layout, int):
return layout
code = _LAYOUT_CODES.get(layout.lower())
if code is None:
raise ValueError(
f"unknown layout '{layout}', expected one of {list(_LAYOUT_CODES)}"
)
return code
def _to_bhld(t: torch.Tensor, layout: int) -> torch.Tensor:
"""Normalize to b h l d view. Zero-copy transpose if layout==1 (b l h d)."""
if layout == 1:
return t.transpose(1, 2)
return t
def _expand_kv_heads(
k: torch.Tensor, v: torch.Tensor, q_head: int
) -> tuple[torch.Tensor, torch.Tensor]:
"""Expand K/V heads to match Q heads for GQA fallback."""
kv_head = k.size(1)
if kv_head == q_head:
return k, v
group = q_head // kv_head
k = k.repeat_interleave(group, dim=1)
v = v.repeat_interleave(group, dim=1)
return k, v
def _build_attn_mask(
q: torch.Tensor,
k: torch.Tensor,
mask: torch.Tensor | None,
causal_offset: int,
scale: float,
) -> tuple[torch.Tensor | None, float]:
"""Build SDPA-compatible attn_mask + resolved scale.
q and k must already be in b h l d layout.
Causal and mask can coexist: causal sets -inf above the diagonal, mask
sets -inf for padded positions. Both are OR'd into a single bool mask.
"""
q_len = q.size(2)
kv_len = k.size(2)
head_dim = q.size(3)
resolved_scale = scale if scale and scale > 0 else 1.0 / math.sqrt(head_dim)
attn_mask = None
if mask is not None:
if mask.dim() == 2:
# [batch, kv_len] → [batch, 1, 1, kv_len]
attn_mask = mask[:, None, None, :]
elif mask.dim() == 3:
# [batch, q_len, kv_len] → [batch, 1, q_len, kv_len]
attn_mask = mask[:, None, :, :]
else:
raise ValueError(f"mask must be 2D or 3D, got {mask.dim()}D")
if causal_offset >= 0:
batch = q.size(0)
# q row i attends to kv cols 0..(causal_offset + i)
q_idx = torch.arange(q_len, device=q.device).unsqueeze(1) # [q_len, 1]
kv_idx = torch.arange(kv_len, device=q.device).unsqueeze(0) # [1, kv_len]
causal_bool = kv_idx > (causal_offset + q_idx) # True = masked out
causal_mask = causal_bool.unsqueeze(0).expand(
batch, -1, -1
) # [batch, q_len, kv_len]
causal_mask = causal_mask[:, None, :, :] # [batch, 1, q_len, kv_len]
if attn_mask is not None:
attn_mask = attn_mask | causal_mask
else:
attn_mask = causal_mask
return attn_mask, resolved_scale
def _torch_fallback(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
mask: torch.Tensor | None,
causal_offset: int,
scale: float,
q_layout: int,
kv_layout: int | None = None,
) -> torch.Tensor:
"""Reference attention via ``scaled_dot_product_attention``.
q_layout / kv_layout: 0 = b h l d, 1 = b l h d.
If kv_layout is None, uses q_layout (Q and K/V share the same layout).
"""
if kv_layout is None:
kv_layout = q_layout
q = _to_bhld(q, q_layout)
k = _to_bhld(k, kv_layout)
v = _to_bhld(v, kv_layout)
k, v = _expand_kv_heads(k, v, q.size(1))
attn_mask, resolved_scale = _build_attn_mask(q, k, mask, causal_offset, scale)
out = F.scaled_dot_product_attention(
q, k, v, attn_mask=attn_mask, is_causal=False, scale=resolved_scale
)
# Restore Q's original layout
if q_layout == 1:
out = out.transpose(1, 2)
return out
def _gather_kv_from_pages(
page_table: torch.Tensor,
k_cache: torch.Tensor,
v_cache: torch.Tensor,
page_size: int,
kv_len: int,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Gather contiguous K/V from paged cache for torch SDPA fallback.
Shapes:
page_table : [batch, max_pages] (int64)
k_cache : [n_pages, page_size, n_kv_heads, head_dim]
v_cache : same as k_cache
Returns:
k, v : [batch, kv_len, n_kv_heads, head_dim] (b l h d)
"""
batch, max_pages = page_table.shape
_, ps, n_kv_heads, head_dim = k_cache.shape
if ps != page_size:
raise ValueError(f"k_cache page_size mismatch: {ps} vs {page_size}")
# Vectorized gather: build physical page + offset indices, then advanced-index
positions = torch.arange(kv_len, device=page_table.device)
logical_pages = positions // page_size # [kv_len]
page_offsets = positions % page_size # [kv_len]
phys_pages = page_table[:, logical_pages] # [batch, kv_len]
# k_cache[phys_pages, page_offsets] → [batch, kv_len, n_kv_heads, head_dim] (b l h d)
k = k_cache[phys_pages, page_offsets]
v = v_cache[phys_pages, page_offsets]
return k, v
def attn_decode(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
mask: torch.Tensor | None = None,
causal_offset: int = -1,
scale: float = 0.0,
layout: str = "bhld",
) -> torch.Tensor:
li = _parse_layout(layout)
if _available["attn_decode"]:
return _modules["attn_decode"].attn_decode(
q,
k,
v,
mask=mask,
causal_offset=causal_offset,
scale=scale,
layout=li,
)
return _torch_fallback(q, k, v, mask, causal_offset, scale, q_layout=li)
def attn_prefill(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
mask: torch.Tensor | None = None,
causal_offset: int = -1,
scale: float = 0.0,
layout: str = "bhld",
) -> torch.Tensor:
li = _parse_layout(layout)
if _available["attn_prefill"]:
return _modules["attn_prefill"].attn_prefill(
q,
k,
v,
mask=mask,
causal_offset=causal_offset,
scale=scale,
layout=li,
)
return _torch_fallback(q, k, v, mask, causal_offset, scale, q_layout=li)
def attn_paged_decode(
q: torch.Tensor,
page_table: torch.Tensor,
k_cache: torch.Tensor,
v_cache: torch.Tensor,
page_size: int,
kv_len: int,
mask: torch.Tensor | None = None,
causal_offset: int = -1,
scale: float = 0.0,
layout: str = "bhld",
) -> torch.Tensor:
li = _parse_layout(layout)
if _available["attn_paged_decode"]:
return _modules["attn_paged_decode"].attn_paged_decode(
q,
page_table,
k_cache,
v_cache,
page_size,
kv_len,
mask=mask,
causal_offset=causal_offset,
scale=scale,
layout=li,
)
# Gathered K/V are always b l h d
k, v = _gather_kv_from_pages(page_table, k_cache, v_cache, page_size, kv_len)
return _torch_fallback(
q, k, v, mask, causal_offset, scale, q_layout=li, kv_layout=1
)
+1 -2
View File
@@ -4,7 +4,6 @@ import inspect
import sys import sys
from abc import ABC from abc import ABC
from typing import ( from typing import (
Any,
Callable, Callable,
Dict, Dict,
ForwardRef, ForwardRef,
@@ -38,7 +37,7 @@ def _resolve_type(
ns = vars(mod) ns = vars(mod)
if isinstance(arg, ForwardRef): if isinstance(arg, ForwardRef):
return arg._evaluate(ns, None, frozenset(), recursive_guard=frozenset()) return arg._evaluate(ns, None, recursive_guard=frozenset())
return ns.get(name) return ns.get(name)
+13 -3
View File
@@ -6,7 +6,7 @@ Layers:
- protocols/: Response builders (OpenAI, Anthropic) - protocols/: Response builders (OpenAI, Anthropic)
- transport/: SSE transport utilities - transport/: SSE transport utilities
- engine.py: Facade (InferenceEngine), Value Object (GenerationRequest) - engine.py: Facade (InferenceEngine), Value Object (GenerationRequest)
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy) - sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy, FrequencyPenaltyStrategy)
""" """
from astrai.inference.api import ( from astrai.inference.api import (
@@ -30,10 +30,14 @@ from astrai.inference.api.openai import OpenAIResponseBuilder
from astrai.inference.core import ( from astrai.inference.core import (
STOP, STOP,
Allocator, Allocator,
CacheView,
ContiguousCache,
ContiguousCacheView,
Executor, Executor,
InferenceScheduler, InferenceScheduler,
KVCache, KVCache,
KvcacheView, PageCache,
PageCacheView,
PagePool, PagePool,
PrefixCache, PrefixCache,
Storage, Storage,
@@ -46,6 +50,7 @@ from astrai.inference.core import (
from astrai.inference.engine import GenerationRequest, InferenceEngine from astrai.inference.engine import GenerationRequest, InferenceEngine
from astrai.inference.sample import ( from astrai.inference.sample import (
BaseSamplingStrategy, BaseSamplingStrategy,
FrequencyPenaltyStrategy,
SamplingPipeline, SamplingPipeline,
TemperatureStrategy, TemperatureStrategy,
TopKStrategy, TopKStrategy,
@@ -63,8 +68,12 @@ __all__ = [
"TaskManager", "TaskManager",
"TaskStatus", "TaskStatus",
"Allocator", "Allocator",
"CacheView",
"KVCache", "KVCache",
"KvcacheView", "ContiguousCache",
"ContiguousCacheView",
"PageCache",
"PageCacheView",
"PagePool", "PagePool",
"PrefixCache", "PrefixCache",
"Storage", "Storage",
@@ -75,6 +84,7 @@ __all__ = [
"TemperatureStrategy", "TemperatureStrategy",
"TopKStrategy", "TopKStrategy",
"TopPStrategy", "TopPStrategy",
"FrequencyPenaltyStrategy",
"SamplingPipeline", "SamplingPipeline",
"ProtocolHandler", "ProtocolHandler",
"StopChecker", "StopChecker",
-1
View File
@@ -21,7 +21,6 @@ logger = logging.getLogger(__name__)
_UNSUPPORTED_PARAMS = ( _UNSUPPORTED_PARAMS = (
"n", "n",
"presence_penalty", "presence_penalty",
"frequency_penalty",
"logit_bias", "logit_bias",
"user", "user",
) )
+1
View File
@@ -125,6 +125,7 @@ class ProtocolHandler:
temperature=self.request.temperature, temperature=self.request.temperature,
top_p=self.request.top_p, top_p=self.request.top_p,
top_k=self.request.top_k, top_k=self.request.top_k,
frequency_penalty=getattr(self.request, "frequency_penalty", 0.0),
) )
if self.request.stream: if self.request.stream:
+25 -6
View File
@@ -7,6 +7,7 @@ Subclasses may optionally consume ``token_ids`` for token-level parsing
(e.g. Harmony / VLM-style parsers). (e.g. Harmony / VLM-style parsers).
""" """
import json
import re import re
import uuid import uuid
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
@@ -117,6 +118,29 @@ def _parse_tool_call_json(json_str: str, complete: bool):
Returns ``(name, args, valid)``. Returns ``(name, args, valid)``.
""" """
if complete:
try:
obj = json.loads(json_str)
except json.JSONDecodeError:
return None, "", False
name = obj.get("name")
if not isinstance(name, str) or not name:
return None, "", False
args = obj.get("arguments")
if isinstance(args, dict):
if not args:
args = ""
else:
args = json.dumps(args, ensure_ascii=False)
args = args[1:-1].rstrip()
elif isinstance(args, list):
args = json.dumps(args, ensure_ascii=False) if args else ""
elif isinstance(args, str):
pass
else:
args = str(args) if args is not None else ""
return name, args, True
name_match = re.search(r'"name"\s*:\s*"([^"]*)"', json_str) name_match = re.search(r'"name"\s*:\s*"([^"]*)"', json_str)
if not name_match: if not name_match:
return None, "", False return None, "", False
@@ -127,8 +151,6 @@ def _parse_tool_call_json(json_str: str, complete: bool):
return name, "", True return name, "", True
raw = args_match.group(1).rstrip() raw = args_match.group(1).rstrip()
if complete and raw.endswith("}"):
raw = raw[:-1].rstrip()
if raw.startswith("{"): if raw.startswith("{"):
inner = raw[1:].rstrip() inner = raw[1:].rstrip()
if inner.endswith("}"): if inner.endswith("}"):
@@ -156,9 +178,6 @@ def _find_tool_calls(text: str, start_pos: int = 0):
break break
json_str = text[brace:end] json_str = text[brace:end]
if not _TOOL_CALL_HEAD_RE.search(json_str):
pos = end
continue
name, args, valid = _parse_tool_call_json(json_str, complete=True) name, args, valid = _parse_tool_call_json(json_str, complete=True)
if not valid or name is None: if not valid or name is None:
@@ -186,7 +205,7 @@ def _find_partial_tool_call(text: str, start_pos: int = 0):
return None return None
json_str = text[brace:] json_str = text[brace:]
if not _TOOL_CALL_HEAD_RE.search(json_str): if '"name"' not in json_str:
return None return None
name, args, valid = _parse_tool_call_json(json_str, complete=False) name, args, valid = _parse_tool_call_json(json_str, complete=False)
+10 -2
View File
@@ -2,8 +2,12 @@
from astrai.inference.core.cache import ( from astrai.inference.core.cache import (
Allocator, Allocator,
CacheView,
ContiguousCache,
ContiguousCacheView,
KVCache, KVCache,
KvcacheView, PageCache,
PageCacheView,
PagePool, PagePool,
PrefixCache, PrefixCache,
Storage, Storage,
@@ -16,8 +20,12 @@ from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
__all__ = [ __all__ = [
"Allocator", "Allocator",
"CacheView",
"KVCache", "KVCache",
"KvcacheView", "ContiguousCache",
"ContiguousCacheView",
"PageCache",
"PageCacheView",
"PagePool", "PagePool",
"PrefixCache", "PrefixCache",
"Storage", "Storage",
+172 -7
View File
@@ -1,4 +1,5 @@
import threading import threading
from abc import ABC, abstractmethod
from collections import OrderedDict from collections import OrderedDict
from typing import Callable, Dict, List, Optional, Tuple from typing import Callable, Dict, List, Optional, Tuple
@@ -62,6 +63,7 @@ class Allocator:
def touch(self, idx: int): def touch(self, idx: int):
with self._lock: with self._lock:
if idx in self._lru:
self._lru.move_to_end(idx) self._lru.move_to_end(idx)
@@ -274,7 +276,46 @@ class Storage:
return k, v return k, v
class KvcacheView: class CacheView(ABC):
"""Abstract view passed to attention layers for KV-cache I/O."""
@abstractmethod
def write(self, layer_id: int, k: Tensor, v: Tensor): ...
@abstractmethod
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]: ...
class KVCache(ABC):
"""Abstract KV-cache facade for scheduler/executor."""
@abstractmethod
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool: ...
@abstractmethod
def task_free(self, task_id: str): ...
@abstractmethod
def task_extend(self, task_id: str, pos: int) -> bool: ...
@abstractmethod
def bind_tasks(
self,
task_ids: List[str],
total_len: int,
device: torch.device,
write_positions: Optional[Tensor] = None,
) -> CacheView: ...
def task_cached(self, task_id: str) -> int:
return 0
def task_record_hashes(
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
): ...
class PageCacheView(CacheView):
"""Bundles Storage + page_table + total_len for attention layers.""" """Bundles Storage + page_table + total_len for attention layers."""
def __init__(self, storage: Storage, page_table: Tensor, total_len: int = 0): def __init__(self, storage: Storage, page_table: Tensor, total_len: int = 0):
@@ -290,8 +331,8 @@ class KvcacheView:
return self._storage.gather(layer_id, self._page_table, self._total_len) return self._storage.gather(layer_id, self._page_table, self._total_len)
class KVCache: class PageCache(KVCache):
"""Facade: page management + KV-cache I/O for continuous batching.""" """Paged KV-cache with prefix sharing."""
def __init__( def __init__(
self, self,
@@ -361,8 +402,132 @@ class KVCache:
for i in range(start_logical_page, full_pages): for i in range(start_logical_page, full_pages):
self._pool.record(page_table[i], prompt_ids, i) self._pool.record(page_table[i], prompt_ids, i)
def make_table_tensor(self, task_ids: List[str], device: torch.device) -> Tensor: def bind_tasks(
return self._table.table_tensor(task_ids, device) self,
task_ids: List[str],
total_len: int,
device: torch.device,
write_positions: Optional[Tensor] = None,
) -> PageCacheView:
page_table = self._table.table_tensor(task_ids, device)
return PageCacheView(self._storage, page_table, total_len)
def bind(self, page_table: Tensor, total_len: int = 0) -> KvcacheView:
return KvcacheView(self._storage, page_table, total_len) class ContiguousCacheView(CacheView):
"""Contiguous KV-cache view for attention layers."""
def __init__(
self,
cache: "ContiguousCache",
batch_indices: Tensor,
total_len: int = 0,
write_positions: Optional[Tensor] = None,
):
self._cache = cache
self._batch_indices = batch_indices
self._total_len = total_len
self._write_positions = write_positions
def write(self, layer_id: int, k: Tensor, v: Tensor):
seq_len = k.size(1)
indices = self._batch_indices
if self._write_positions is not None and seq_len == 1:
pos = self._write_positions
self._cache.k[layer_id, indices, pos] = k.squeeze(1)
self._cache.v[layer_id, indices, pos] = v.squeeze(1)
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()
)
indices = self._batch_indices
k = self._cache.k[layer_id, indices, :max_len]
v = self._cache.v[layer_id, indices, :max_len]
return k, v
class ContiguousCache(KVCache):
"""Contiguous per-slot KV cache (default implementation)."""
def __init__(
self,
n_layers: int,
max_batch_size: int,
max_seq_len: int,
n_kv_heads: int,
head_dim: int,
device: torch.device,
dtype: torch.dtype,
):
self.max_seq_len = max_seq_len
self.k = torch.zeros(
n_layers,
max_batch_size,
max_seq_len,
n_kv_heads,
head_dim,
device=device,
dtype=dtype,
)
self.v = torch.zeros(
n_layers,
max_batch_size,
max_seq_len,
n_kv_heads,
head_dim,
device=device,
dtype=dtype,
)
self._slot_len: Dict[int, int] = {}
self._task_slot: Dict[str, int] = {}
self._free_slots = list(range(max_batch_size))
self._device = device
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
if not self._free_slots:
return False
slot = self._free_slots.pop(0)
self._task_slot[task_id] = slot
self._slot_len[slot] = 0
return True
def task_free(self, task_id: str):
slot = self._task_slot.pop(task_id, None)
if slot is not None:
self._slot_len.pop(slot, None)
self._free_slots.append(slot)
def task_extend(self, task_id: str, pos: int) -> bool:
return pos < self.max_seq_len
def task_cached(self, task_id: str) -> int:
slot = self._task_slot.get(task_id)
if slot is None:
return 0
return self._slot_len.get(slot, 0)
def bind_tasks(
self,
task_ids: List[str],
total_len: int,
device: torch.device,
write_positions: Optional[Tensor] = None,
) -> ContiguousCacheView:
slots = [self._task_slot[tid] for tid in task_ids]
batch_indices = torch.tensor(slots, dtype=torch.long, device=device)
return ContiguousCacheView(
self, batch_indices, total_len, write_positions=write_positions
)
+39 -6
View File
@@ -19,13 +19,13 @@ class Executor:
self, self,
model: AutoModel, model: AutoModel,
tokenizer: AutoTokenizer, tokenizer: AutoTokenizer,
page_cache: KVCache, kv_cache: KVCache,
device: Optional[str] = None, device: Optional[str] = None,
dtype: Optional[torch.dtype] = None, dtype: Optional[torch.dtype] = None,
): ):
self.model = model self.model = model
self.tokenizer = tokenizer self.tokenizer = tokenizer
self.page_cache = page_cache self.kv_cache = kv_cache
self.device = device or next(model.parameters()).device self.device = device or next(model.parameters()).device
self.dtype = dtype or next(model.parameters()).dtype self.dtype = dtype or next(model.parameters()).dtype
@@ -43,7 +43,6 @@ class Executor:
) )
task_ids = [t.task_id for t in tasks] task_ids = [t.task_id for t in tasks]
page_tables = self.page_cache.make_table_tensor(task_ids, self.device)
with torch.inference_mode(): with torch.inference_mode():
self.model( self.model(
@@ -53,7 +52,7 @@ class Executor:
) )
.unsqueeze(0) .unsqueeze(0)
.expand(batch_sz, -1), .expand(batch_sz, -1),
paged_cache=self.page_cache.bind(page_tables, total_len=prompt_len), 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]) -> List[int]:
@@ -72,16 +71,47 @@ class Executor:
total_len = position_ids.max().item() + 1 total_len = position_ids.max().item() + 1
task_ids = [t.task_id for t in tasks] task_ids = [t.task_id for t in tasks]
page_tables = self.page_cache.make_table_tensor(task_ids, self.device)
temperatures = torch.tensor([t.temperature for t in tasks], device=self.device) temperatures = torch.tensor([t.temperature for t in tasks], device=self.device)
top_ks = torch.tensor([t.top_k for t in tasks], device=self.device) top_ks = torch.tensor([t.top_k for t in tasks], device=self.device)
top_ps = torch.tensor([t.top_p for t in tasks], device=self.device) top_ps = torch.tensor([t.top_p for t in tasks], device=self.device)
freq_penalties = torch.tensor(
[t.frequency_penalty for t in tasks], device=self.device
)
history_lists = []
mask_lists = []
for t in tasks:
window = t.rep_window
prompt_part = t.prompt_ids[-window:]
ids = prompt_part + t.output_ids
history_lists.append(ids)
mask_lists.append([True] * len(ids))
max_len = max(len(h) for h in history_lists)
padded_ids = torch.zeros(
len(tasks), max_len, dtype=torch.long, device=self.device
)
padded_mask = torch.zeros(
len(tasks), max_len, dtype=torch.bool, device=self.device
)
for i, (h, m) in enumerate(zip(history_lists, mask_lists)):
padded_ids[i, : len(h)] = torch.tensor(
h, dtype=torch.long, device=self.device
)
padded_mask[i, : len(m)] = torch.tensor(
m, dtype=torch.bool, device=self.device
)
with torch.inference_mode(): with torch.inference_mode():
outputs = self.model( outputs = self.model(
input_ids.unsqueeze(1), input_ids.unsqueeze(1),
paged_cache=self.page_cache.bind(page_tables, total_len=total_len), paged_cache=self.kv_cache.bind_tasks(
task_ids,
total_len,
self.device,
write_positions=position_ids,
),
position_ids=position_ids.unsqueeze(1), position_ids=position_ids.unsqueeze(1),
) )
logits = outputs["logits"][:, -1, :] logits = outputs["logits"][:, -1, :]
@@ -91,4 +121,7 @@ class Executor:
temperature=temperatures, temperature=temperatures,
top_k=top_ks, top_k=top_ks,
top_p=top_ps, top_p=top_ps,
frequency_penalty=freq_penalties,
input_ids=padded_ids,
input_mask=padded_mask,
).tolist() ).tolist()
+40 -54
View File
@@ -4,7 +4,7 @@ from typing import Any, Dict, List, Optional, Tuple
import torch import torch
from astrai.inference.core.cache import KVCache from astrai.inference.core.cache import ContiguousCache, KVCache
from astrai.inference.core.executor import Executor from astrai.inference.core.executor import Executor
from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
from astrai.model.automodel import AutoModel from astrai.model.automodel import AutoModel
@@ -14,7 +14,7 @@ logger = logging.getLogger(__name__)
class InferenceScheduler: class InferenceScheduler:
"""Four-phase continuous batching loop: cleanup -> refill -> prefill -> decode.""" """Continuous batching loop: cleanup -> refill -> prefill -> decode (all groups)."""
def __init__( def __init__(
self, self,
@@ -23,9 +23,9 @@ class InferenceScheduler:
max_batch_size: int = 16, max_batch_size: int = 16,
max_seq_len: Optional[int] = None, max_seq_len: Optional[int] = None,
max_prompt_len: int = 2048, max_prompt_len: int = 2048,
page_size: int = 64,
device: Optional[str] = None, device: Optional[str] = None,
dtype: Optional[torch.dtype] = None, dtype: Optional[torch.dtype] = None,
cache: Optional[KVCache] = None,
): ):
config = model.config config = model.config
@@ -41,16 +41,17 @@ class InferenceScheduler:
self.device = device or next(model.parameters()).device self.device = device or next(model.parameters()).device
self.dtype = dtype or next(model.parameters()).dtype self.dtype = dtype or next(model.parameters()).dtype
n_pages = ( head_dim = config.dim // config.n_heads
max_batch_size * (self.max_seq_len + page_size) + page_size - 1
) // page_size
self._page_cache = KVCache( if cache is not None:
self._cache = cache
else:
self._cache = ContiguousCache(
config.n_layers, config.n_layers,
n_pages, max_batch_size,
page_size, self.max_seq_len,
config.n_kv_heads, config.n_kv_heads,
config.dim // config.n_heads, head_dim,
self.device, self.device,
self.dtype, self.dtype,
) )
@@ -65,7 +66,7 @@ class InferenceScheduler:
self._executor = Executor( self._executor = Executor(
model=model, model=model,
tokenizer=tokenizer, tokenizer=tokenizer,
page_cache=self._page_cache, kv_cache=self._cache,
device=self.device, device=self.device,
dtype=self.dtype, dtype=self.dtype,
) )
@@ -78,18 +79,19 @@ class InferenceScheduler:
def remove_task(self, task_id: str): def remove_task(self, task_id: str):
for task in self._task_mgr.remove_task(task_id): for task in self._task_mgr.remove_task(task_id):
self._page_cache.task_free(task.task_id) self._cache.task_free(task.task_id)
def get_stats(self) -> Dict[str, Any]: def get_stats(self) -> Dict[str, Any]:
return self._task_mgr.get_stats() return self._task_mgr.get_stats()
def _run_generation_loop(self): def _run_generation_loop(self):
stop_ids = self._task_mgr.tokenizer.stop_ids stop_ids = self._task_mgr.tokenizer.stop_ids
cache = self._cache
try: try:
while not self._stop_event.is_set(): while not self._stop_event.is_set():
finished = self._task_mgr.remove_finished_tasks(stop_ids) finished = self._task_mgr.remove_finished_tasks(stop_ids)
for task in finished: for task in finished:
self._page_cache.task_free(task.task_id) cache.task_free(task.task_id)
active = self._task_mgr.get_active_tasks() active = self._task_mgr.get_active_tasks()
available = self._task_mgr.max_batch_size - len(active) available = self._task_mgr.max_batch_size - len(active)
@@ -97,7 +99,7 @@ class InferenceScheduler:
candidates = self._task_mgr.pull_candidates(available) candidates = self._task_mgr.pull_candidates(available)
failed = [] failed = []
for task in candidates: for task in candidates:
if self._page_cache.task_alloc(task.task_id, task.prompt_ids): if cache.task_alloc(task.task_id, task.prompt_ids):
self._task_mgr.activate(task) self._task_mgr.activate(task)
else: else:
failed.append(task) failed.append(task)
@@ -112,7 +114,7 @@ class InferenceScheduler:
t t
for t in self._task_mgr.get_active_tasks() for t in self._task_mgr.get_active_tasks()
if t.output_tokens == 0 if t.output_tokens == 0
and self._page_cache.task_cached(t.task_id) < len(t.prompt_ids) and cache.task_cached(t.task_id) < len(t.prompt_ids)
] ]
if to_prefill: if to_prefill:
for t in to_prefill: for t in to_prefill:
@@ -122,36 +124,29 @@ class InferenceScheduler:
for t in to_prefill: for t in to_prefill:
key = ( key = (
len(t.prompt_ids), len(t.prompt_ids),
self._page_cache.task_cached(t.task_id), cache.task_cached(t.task_id),
) )
groups.setdefault(key, []).append(t) groups.setdefault(key, []).append(t)
for (prompt_len, start_pos), group in groups.items(): for (prompt_len, start_pos), group in groups.items():
self._executor.execute_prefill(group, prompt_len, start_pos) self._executor.execute_prefill(group, prompt_len, start_pos)
start_logical_page = start_pos // self._page_cache.page_size start_logical_page = start_pos // getattr(
cache, "page_size", 64
)
for t in group: for t in group:
self._page_cache.task_record_hashes( cache.task_record_hashes(
t.task_id, t.task_id, t.prompt_ids, start_logical_page
t.prompt_ids,
start_logical_page=start_logical_page,
) )
pos_groups: Dict[int, List[Task]] = {} decode_tasks = self._task_mgr.get_active_tasks()
for t in self._task_mgr.get_active_tasks():
pos_groups.setdefault(t.next_pos, []).append(t)
if pos_groups:
best_key = max(pos_groups, key=lambda k: len(pos_groups[k]))
group = sorted(pos_groups[best_key], key=lambda t: t.task_id)
valid: List[Task] = [] valid: List[Task] = []
for t in group: for t in sorted(decode_tasks, key=lambda t: t.task_id):
if self._page_cache.task_extend(t.task_id, t.next_pos): if cache.task_extend(t.task_id, t.next_pos):
valid.append(t) valid.append(t)
else: else:
t.status = TaskStatus.ABORTED t.status = TaskStatus.ABORTED
if t.stream_callback: self._task_mgr.invoke_callback(t.task_id, STOP)
t.stream_callback(STOP)
if valid: if valid:
next_tokens = self._executor.execute_decode(valid) next_tokens = self._executor.execute_decode(valid)
@@ -159,32 +154,25 @@ class InferenceScheduler:
for t, ntok in zip(valid, next_tokens): for t, ntok in zip(valid, next_tokens):
t.output_ids.append(ntok) t.output_ids.append(ntok)
t.output_tokens += 1 t.output_tokens += 1
pos = t.input_tokens + t.output_tokens new_text = t.decode_new_token(self._task_mgr.tokenizer)
extend_ok = self._page_cache.task_extend(t.task_id, pos) if new_text:
if t.stream_callback: self._task_mgr.invoke_callback(t.task_id, new_text)
t.stream_callback(
self._task_mgr.tokenizer.decode([ntok])
)
if not extend_ok:
t.status = TaskStatus.ABORTED
if t.stream_callback:
t.stream_callback(STOP)
for t in valid: for t in valid:
if t.is_finished(stop_ids): if t.is_finished(stop_ids):
if t.stream_callback: remaining = t.flush_remaining(self._task_mgr.tokenizer)
t.stream_callback(STOP) if remaining:
self._task_mgr.invoke_callback(t.task_id, remaining)
self._task_mgr.invoke_callback(t.task_id, STOP)
except Exception as e: except Exception as e:
self._stop_event.set() self._stop_event.set()
logger.error(f"Scheduler loop crashed: {e}", exc_info=True) logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
for task in self._task_mgr.get_active_tasks(): for task in self._task_mgr.get_active_tasks():
if task.stream_callback: self._task_mgr.invoke_callback(task.task_id, STOP)
task.stream_callback(STOP) cache.task_free(task.task_id)
self._page_cache.task_free(task.task_id)
for task in self._task_mgr.get_waiting_tasks(): for task in self._task_mgr.get_waiting_tasks():
if task.stream_callback: self._task_mgr.invoke_callback(task.task_id, STOP)
task.stream_callback(STOP)
self._task_mgr.clear_queues() self._task_mgr.clear_queues()
def start(self): def start(self):
@@ -202,12 +190,10 @@ class InferenceScheduler:
self._loop_thread.join(timeout=2.0) self._loop_thread.join(timeout=2.0)
self._loop_thread = None self._loop_thread = None
for task in self._task_mgr.get_active_tasks(): for task in self._task_mgr.get_active_tasks():
if task.stream_callback: self._task_mgr.invoke_callback(task.task_id, STOP)
task.stream_callback(STOP) self._cache.task_free(task.task_id)
self._page_cache.task_free(task.task_id)
for task in self._task_mgr.get_waiting_tasks(): for task in self._task_mgr.get_waiting_tasks():
if task.stream_callback: self._task_mgr.invoke_callback(task.task_id, STOP)
task.stream_callback(STOP)
self._task_mgr.clear_queues() self._task_mgr.clear_queues()
if torch.cuda.is_available(): if torch.cuda.is_available():
torch.cuda.empty_cache() torch.cuda.empty_cache()
+80 -3
View File
@@ -13,6 +13,40 @@ logger = logging.getLogger(__name__)
STOP = object() STOP = object()
class StreamDecoder:
"""Incremental decoder for byte-level BPE streaming.
Byte-level BPE may split a single Unicode character (e.g. em-dash,
smart quotes) across multiple tokens. Decoding such a token in
isolation produces U+FFFD (replacement char). This decoder
accumulates token IDs and only emits text once the trailing
characters are complete, buffering incomplete multi-byte sequences
until the next token arrives.
"""
__slots__ = ("_tokenizer", "_ids", "_emitted")
def __init__(self, tokenizer: AutoTokenizer):
self._tokenizer = tokenizer
self._ids: List[int] = []
self._emitted: str = ""
def push(self, token_id: int) -> str:
"""Append a token ID and return newly completed text.
Returns "" while a multi-byte character is still incomplete.
"""
self._ids.append(token_id)
full = self._tokenizer.decode(self._ids, skip_special_tokens=True)
if full.endswith("\ufffd"):
return ""
if len(full) > len(self._emitted):
diff = full[len(self._emitted) :]
self._emitted = full
return diff
return ""
class TaskStatus(Enum): class TaskStatus(Enum):
"""Task lifecycle states.""" """Task lifecycle states."""
@@ -33,7 +67,8 @@ class Task:
temperature: float = 1.0, temperature: float = 1.0,
top_p: float = 1.0, top_p: float = 1.0,
top_k: int = 50, top_k: int = 50,
stream_callback: Optional[Callable[[str], None]] = None, frequency_penalty: float = 0.0,
rep_window: int = 64,
): ):
self.task_id = task_id self.task_id = task_id
self.prompt_ids = prompt_ids self.prompt_ids = prompt_ids
@@ -41,6 +76,8 @@ class Task:
self.temperature = temperature self.temperature = temperature
self.top_p = top_p self.top_p = top_p
self.top_k = top_k self.top_k = top_k
self.frequency_penalty = frequency_penalty
self.rep_window = rep_window
self.status = TaskStatus.PENDING self.status = TaskStatus.PENDING
self.output_ids: List[int] = [] self.output_ids: List[int] = []
@@ -48,7 +85,34 @@ class Task:
self.output_tokens: int = 0 self.output_tokens: int = 0
self.arrival_time = time.time() self.arrival_time = time.time()
self.finish_time: Optional[float] = None self.finish_time: Optional[float] = None
self.stream_callback = stream_callback self._decoder: Optional[StreamDecoder] = None
def decode_new_token(self, tokenizer: AutoTokenizer) -> str:
"""Decode the last appended output token, buffering incomplete
multi-byte sequences across calls.
Lazily creates a :class:`StreamDecoder` on first use.
"""
if self._decoder is None:
self._decoder = StreamDecoder(tokenizer)
return self._decoder.push(self.output_ids[-1])
def flush_remaining(self, tokenizer: AutoTokenizer) -> str:
"""Emit any text still buffered in the decoder.
Called when generation terminates (max_tokens reached, stop
sequence, or external removal) to avoid dropping a final
incomplete-looking fragment that is actually complete when
adjacent to the stop token.
"""
if self._decoder is None or not self.output_ids:
return ""
full = tokenizer.decode(self.output_ids, skip_special_tokens=True)
if len(full) > len(self._decoder._emitted):
diff = full[len(self._decoder._emitted) :]
self._decoder._emitted = full
return diff
return ""
@property @property
def next_pos(self) -> int: def next_pos(self) -> int:
@@ -79,6 +143,7 @@ class TaskManager:
self.waiting_queue: Deque[Task] = deque() self.waiting_queue: Deque[Task] = deque()
self.active_tasks: List[Task] = [] self.active_tasks: List[Task] = []
self._callbacks: Dict[str, Callable[[str], None]] = {}
self._task_event = threading.Event() self._task_event = threading.Event()
self._lock = threading.Lock() self._lock = threading.Lock()
@@ -93,6 +158,8 @@ class TaskManager:
temperature: float = 1.0, temperature: float = 1.0,
top_p: float = 1.0, top_p: float = 1.0,
top_k: int = 50, top_k: int = 50,
frequency_penalty: float = 0.0,
rep_window: int = 64,
stream_callback: Optional[Callable[[str], None]] = None, stream_callback: Optional[Callable[[str], None]] = None,
) -> str: ) -> str:
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}" task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
@@ -117,12 +184,15 @@ class TaskManager:
temperature=temperature, temperature=temperature,
top_p=top_p, top_p=top_p,
top_k=top_k, top_k=top_k,
stream_callback=stream_callback, frequency_penalty=frequency_penalty,
rep_window=rep_window,
) )
with self._lock: with self._lock:
self.waiting_queue.append(task) self.waiting_queue.append(task)
self._total_tasks += 1 self._total_tasks += 1
if stream_callback:
self._callbacks[task_id] = stream_callback
self._task_event.set() self._task_event.set()
return task_id return task_id
@@ -134,8 +204,14 @@ class TaskManager:
t for t in self.waiting_queue if t.task_id != task_id t for t in self.waiting_queue if t.task_id != task_id
) )
self.active_tasks = [t for t in self.active_tasks if t.task_id != task_id] self.active_tasks = [t for t in self.active_tasks if t.task_id != task_id]
self._callbacks.pop(task_id, None)
return removed_active return removed_active
def invoke_callback(self, task_id: str, token: str):
cb = self._callbacks.get(task_id)
if cb:
cb(token)
def get_stats(self) -> Dict[str, Any]: def get_stats(self) -> Dict[str, Any]:
return { return {
"total_tasks": self._total_tasks, "total_tasks": self._total_tasks,
@@ -204,6 +280,7 @@ class TaskManager:
with self._lock: with self._lock:
self.waiting_queue.clear() self.waiting_queue.clear()
self.active_tasks.clear() self.active_tasks.clear()
self._callbacks.clear()
def wake(self): def wake(self):
self._task_event.set() self._task_event.set()
+68 -8
View File
@@ -8,6 +8,7 @@ from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Tuple,
import torch import torch
import torch.nn as nn import torch.nn as nn
from astrai.inference.core.cache import KVCache
from astrai.inference.core.scheduler import InferenceScheduler from astrai.inference.core.scheduler import InferenceScheduler
from astrai.inference.core.task import STOP from astrai.inference.core.task import STOP
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
@@ -73,20 +74,31 @@ class GenerationRequest:
top_p: float = 1.0, top_p: float = 1.0,
temperature: float = 1.0, temperature: float = 1.0,
max_tokens: Optional[int] = None, max_tokens: Optional[int] = None,
frequency_penalty: float = 0.0,
rep_window: int = 64,
stream: bool = False, stream: bool = False,
): ):
if not (isinstance(top_k, int) and top_k >= 0): if not (isinstance(top_k, int) and top_k >= 0):
raise ValueError("top_k must be a non-negative integer") raise ValueError("top_k must be a non-negative integer")
if not (0.0 <= top_p <= 1.0): if not (0.0 <= top_p <= 1.0):
raise ValueError("top_p must be a float between 0.0 and 1.0") raise ValueError("top_p must be a float between 0.0 and 1.0")
if not (isinstance(temperature, (int, float)) and temperature > 0): if not (isinstance(temperature, (int, float)) and temperature >= 0):
raise ValueError("temperature must be a positive number") raise ValueError("temperature must be a non-negative number")
if not (
isinstance(frequency_penalty, (int, float))
and -2.0 <= frequency_penalty <= 2.0
):
raise ValueError("frequency_penalty must be between -2.0 and 2.0")
if not (isinstance(rep_window, int) and rep_window > 0):
raise ValueError("rep_window must be a positive integer")
self.messages = messages self.messages = messages
self.top_k = top_k self.top_k = top_k
self.top_p = top_p self.top_p = top_p
self.temperature = temperature self.temperature = temperature
self.max_tokens = max_tokens self.max_tokens = max_tokens
self.frequency_penalty = frequency_penalty
self.rep_window = rep_window
self.stream = stream self.stream = stream
@@ -101,6 +113,7 @@ class InferenceEngine:
max_seq_len: Optional[int] = None, max_seq_len: Optional[int] = None,
max_prompt_len: int = 2048, max_prompt_len: int = 2048,
page_size: int = 128, page_size: int = 128,
cache: Optional[KVCache] = None,
): ):
self.model = model self.model = model
self.tokenizer = tokenizer self.tokenizer = tokenizer
@@ -110,7 +123,7 @@ class InferenceEngine:
max_batch_size=max_batch_size, max_batch_size=max_batch_size,
max_seq_len=max_seq_len, max_seq_len=max_seq_len,
max_prompt_len=max_prompt_len, max_prompt_len=max_prompt_len,
page_size=page_size, cache=cache,
) )
self.scheduler.start() self.scheduler.start()
@@ -130,17 +143,33 @@ class InferenceEngine:
temperature: float = 1.0, temperature: float = 1.0,
top_p: float = 1.0, top_p: float = 1.0,
top_k: int = 50, top_k: int = 50,
frequency_penalty: float = 0.0,
rep_window: int = 64,
) -> Union[Generator, str, List[str]]: ) -> Union[Generator, str, List[str]]:
is_batch = isinstance(prompt, list) is_batch = isinstance(prompt, list)
prompts = prompt if is_batch else [prompt] prompts = prompt if is_batch else [prompt]
if stream: if stream:
return self._generate_streaming( return self._generate_streaming(
prompts, is_batch, max_tokens, temperature, top_p, top_k prompts,
is_batch,
max_tokens,
temperature,
top_p,
top_k,
frequency_penalty,
rep_window,
) )
else: else:
return self._generate_non_streaming( return self._generate_non_streaming(
prompts, is_batch, max_tokens, temperature, top_p, top_k prompts,
is_batch,
max_tokens,
temperature,
top_p,
top_k,
frequency_penalty,
rep_window,
) )
def generate_async( def generate_async(
@@ -150,9 +179,18 @@ class InferenceEngine:
temperature: float = 1.0, temperature: float = 1.0,
top_p: float = 1.0, top_p: float = 1.0,
top_k: int = 50, top_k: int = 50,
frequency_penalty: float = 0.0,
rep_window: int = 64,
) -> AsyncGenerator[str, None]: ) -> AsyncGenerator[str, None]:
sync_gen = self._generate_streaming( sync_gen = self._generate_streaming(
[prompt], False, max_tokens, temperature, top_p, top_k [prompt],
False,
max_tokens,
temperature,
top_p,
top_k,
frequency_penalty,
rep_window,
) )
async def _agen(): async def _agen():
@@ -183,6 +221,8 @@ class InferenceEngine:
temperature=request.temperature, temperature=request.temperature,
top_p=request.top_p, top_p=request.top_p,
top_k=request.top_k, top_k=request.top_k,
frequency_penalty=request.frequency_penalty,
rep_window=request.rep_window,
) )
def _submit_tasks( def _submit_tasks(
@@ -192,6 +232,8 @@ class InferenceEngine:
temperature: float, temperature: float,
top_p: float, top_p: float,
top_k: int, top_k: int,
frequency_penalty: float,
rep_window: int,
) -> Tuple[GenerateResult, List[str]]: ) -> Tuple[GenerateResult, List[str]]:
n = len(prompts) n = len(prompts)
result = GenerateResult(count=n) result = GenerateResult(count=n)
@@ -204,6 +246,8 @@ class InferenceEngine:
temperature=temperature, temperature=temperature,
top_p=top_p, top_p=top_p,
top_k=top_k, top_k=top_k,
frequency_penalty=frequency_penalty,
rep_window=rep_window,
stream_callback=cb, stream_callback=cb,
) )
task_ids.append(task_id) task_ids.append(task_id)
@@ -224,9 +268,17 @@ class InferenceEngine:
temperature: float, temperature: float,
top_p: float, top_p: float,
top_k: int, top_k: int,
frequency_penalty: float,
rep_window: int,
) -> Generator: ) -> Generator:
result, task_ids = self._submit_tasks( result, task_ids = self._submit_tasks(
prompts, max_tokens, temperature, top_p, top_k prompts,
max_tokens,
temperature,
top_p,
top_k,
frequency_penalty,
rep_window,
) )
n = len(prompts) n = len(prompts)
remaining = n remaining = n
@@ -260,9 +312,17 @@ class InferenceEngine:
temperature: float, temperature: float,
top_p: float, top_p: float,
top_k: int, top_k: int,
frequency_penalty: float,
rep_window: int,
) -> Union[str, List[str]]: ) -> Union[str, List[str]]:
result, task_ids = self._submit_tasks( result, task_ids = self._submit_tasks(
prompts, max_tokens, temperature, top_p, top_k prompts,
max_tokens,
temperature,
top_p,
top_k,
frequency_penalty,
rep_window,
) )
try: try:
+163 -13
View File
@@ -1,15 +1,15 @@
"""Composable sampling strategies for logit transformation. """Composable sampling strategies for logit transformation.
Implements the Strategy pattern: each sampling technique Implements the Strategy pattern: each sampling technique
(temperature, top-k, top-p) is a pluggable strategy that (temperature, top-k, top-p, frequency penalty) is a pluggable
can be composed into a pipeline. strategy that can be composed into a pipeline.
All strategies accept both scalar and per-sample tensor All strategies accept both scalar and per-sample tensor
parameters, so a single pipeline works for any batch size. parameters, so a single pipeline works for any batch size.
""" """
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import List, Union from typing import List, Optional, Union
import torch import torch
from torch import Tensor from torch import Tensor
@@ -19,12 +19,23 @@ class BaseSamplingStrategy(ABC):
"""Abstract base for a logit transformation strategy.""" """Abstract base for a logit transformation strategy."""
@abstractmethod @abstractmethod
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor: def apply(
self,
logits: Tensor,
filter_value: float = -float("inf"),
input_ids: Optional[Tensor] = None,
input_mask: Optional[Tensor] = None,
) -> Tensor:
"""Applies the strategy to logits. """Applies the strategy to logits.
Args: Args:
logits: Raw logits tensor (batch, vocab_size). logits: Raw logits tensor (batch, vocab_size).
filter_value: Value assigned to filtered-out positions. filter_value: Value assigned to filtered-out positions.
input_ids: Previously generated token IDs ``[batch, seq_len]``,
padded with 0. Used by frequency penalty.
input_mask: Boolean mask ``[batch, seq_len]``, True for real
tokens, False for padding. Used to exclude padding from
penalty computation.
Returns: Returns:
Transformed logits tensor. Transformed logits tensor.
@@ -42,7 +53,13 @@ class TemperatureStrategy(BaseSamplingStrategy):
def __init__(self, temperature: Union[float, Tensor] = 1.0): def __init__(self, temperature: Union[float, Tensor] = 1.0):
self.temperature = temperature self.temperature = temperature
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor: def apply(
self,
logits: Tensor,
filter_value: float = -float("inf"),
input_ids: Optional[Tensor] = None,
input_mask: Optional[Tensor] = None,
) -> Tensor:
t = self.temperature t = self.temperature
if isinstance(t, Tensor): if isinstance(t, Tensor):
t = t.to(logits.device, non_blocking=True).view(-1, 1) t = t.to(logits.device, non_blocking=True).view(-1, 1)
@@ -64,7 +81,13 @@ class TopKStrategy(BaseSamplingStrategy):
def __init__(self, top_k: Union[int, Tensor] = 0): def __init__(self, top_k: Union[int, Tensor] = 0):
self.top_k = top_k self.top_k = top_k
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor: def apply(
self,
logits: Tensor,
filter_value: float = -float("inf"),
input_ids: Optional[Tensor] = None,
input_mask: Optional[Tensor] = None,
) -> Tensor:
tk = self.top_k tk = self.top_k
if isinstance(tk, Tensor): if isinstance(tk, Tensor):
tk = tk.to(logits.device, non_blocking=True).long().clamp(min=0) tk = tk.to(logits.device, non_blocking=True).long().clamp(min=0)
@@ -114,7 +137,13 @@ class TopPStrategy(BaseSamplingStrategy):
logits[mask] = filter_value logits[mask] = filter_value
return logits return logits
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor: def apply(
self,
logits: Tensor,
filter_value: float = -float("inf"),
input_ids: Optional[Tensor] = None,
input_mask: Optional[Tensor] = None,
) -> Tensor:
tp = self.top_p tp = self.top_p
if isinstance(tp, Tensor): if isinstance(tp, Tensor):
tp = tp.to(logits.device, non_blocking=True) tp = tp.to(logits.device, non_blocking=True)
@@ -125,6 +154,84 @@ class TopPStrategy(BaseSamplingStrategy):
return logits return logits
class FrequencyPenaltyStrategy(BaseSamplingStrategy):
"""Penalizes tokens based on how many times they appeared in history.
Subtracts ``penalty * count(token)`` from each token's logit, where
``count(token)`` is the number of occurrences in the generation history
(prompt + output). A penalty of ``0.0`` disables the strategy.
Unlike repetition penalty (which only checks *presence*), frequency
penalty scales linearly with occurrence count: the first use is
penalized once, the third use three times. This allows natural
repetition of common words while suppressing degenerate loops.
Reference: OpenAI API ``frequency_penalty`` parameter.
Args:
penalty: Scalar or ``[batch]`` tensor (0.0 disables, range -2.0~2.0).
"""
def __init__(self, penalty: Union[float, Tensor] = 0.0):
self.penalty = penalty
def apply(
self,
logits: Tensor,
filter_value: float = -float("inf"),
input_ids: Optional[Tensor] = None,
input_mask: Optional[Tensor] = None,
) -> Tensor:
if input_ids is None:
return logits
p = self.penalty
if isinstance(p, Tensor):
p = p.to(logits.device, non_blocking=True).view(-1, 1)
if (p == 0.0).all():
return logits
elif p == 0.0:
return logits
input_ids = input_ids.to(logits.device, non_blocking=True)
if input_mask is not None:
input_mask = input_mask.to(logits.device, non_blocking=True)
masked_ids = input_ids.clone()
masked_ids[~input_mask] = -1
else:
masked_ids = input_ids
batch_sz, seq_len = masked_ids.shape
vocab_size = logits.size(-1)
if isinstance(p, Tensor):
penalty_per_row = p.expand(batch_sz, 1)
else:
penalty_per_row = torch.full(
(batch_sz, 1), float(p), device=logits.device, dtype=logits.dtype
)
counts = torch.zeros(
batch_sz, vocab_size, device=logits.device, dtype=logits.dtype
)
valid_mask = masked_ids >= 0
if valid_mask.any():
valid_ids = masked_ids[valid_mask]
row_indices = (
torch.arange(batch_sz, device=logits.device)
.unsqueeze(1)
.expand_as(masked_ids)[valid_mask]
)
counts.index_put_(
(row_indices, valid_ids),
torch.ones_like(valid_ids, dtype=logits.dtype),
accumulate=True,
)
return logits - penalty_per_row * counts
class SamplingPipeline(BaseSamplingStrategy): class SamplingPipeline(BaseSamplingStrategy):
"""Composes multiple sampling strategies into a single transformation. """Composes multiple sampling strategies into a single transformation.
@@ -145,23 +252,53 @@ class SamplingPipeline(BaseSamplingStrategy):
def __init__(self, strategies: List[BaseSamplingStrategy]): def __init__(self, strategies: List[BaseSamplingStrategy]):
self.strategies = strategies self.strategies = strategies
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor: def apply(
self,
logits: Tensor,
filter_value: float = -float("inf"),
input_ids: Optional[Tensor] = None,
input_mask: Optional[Tensor] = None,
) -> Tensor:
for strategy in self.strategies: for strategy in self.strategies:
logits = strategy.apply(logits, filter_value) logits = strategy.apply(logits, filter_value, input_ids, input_mask)
return logits return logits
@torch.no_grad() @staticmethod
def sample(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor: def _is_greedy(temperature: Union[float, Tensor]) -> bool:
if isinstance(temperature, Tensor):
return temperature.numel() == 1 and temperature.item() == 0
return temperature == 0
@torch.inference_mode()
def sample(
self,
logits: Tensor,
filter_value: float = -float("inf"),
input_ids: Optional[Tensor] = None,
input_mask: Optional[Tensor] = None,
) -> Tensor:
"""Apply strategies then sample (softmax + multinomial). """Apply strategies then sample (softmax + multinomial).
Short-circuits to ``argmax`` when temperature is exactly 0
(deterministic / greedy decode).
Args: Args:
logits: Raw logits ``[batch, vocab_size]``. logits: Raw logits ``[batch, vocab_size]``.
input_ids: Previously generated token IDs ``[batch, seq_len]``.
input_mask: Boolean mask for ``input_ids`` padding.
Returns: Returns:
Sampled token IDs ``[batch]``. Sampled token IDs ``[batch]``.
""" """
for s in self.strategies:
if isinstance(s, TemperatureStrategy) and self._is_greedy(s.temperature):
return logits.argmax(dim=-1)
break
return torch.multinomial( return torch.multinomial(
torch.softmax(self.apply(logits, filter_value), dim=-1), torch.softmax(
self.apply(logits, filter_value, input_ids, input_mask), dim=-1
),
num_samples=1, num_samples=1,
).squeeze(-1) ).squeeze(-1)
@@ -172,22 +309,35 @@ def sample(
temperature: Union[float, Tensor] = 1.0, temperature: Union[float, Tensor] = 1.0,
top_k: Union[int, Tensor] = 0, top_k: Union[int, Tensor] = 0,
top_p: Union[float, Tensor] = 1.0, top_p: Union[float, Tensor] = 1.0,
frequency_penalty: Union[float, Tensor] = 0.0,
input_ids: Optional[Tensor] = None,
input_mask: Optional[Tensor] = None,
filter_value: float = -float("inf"), filter_value: float = -float("inf"),
) -> Tensor: ) -> Tensor:
"""Apply sampling strategies then sample (softmax + multinomial). """Apply sampling strategies then sample (softmax + multinomial).
Shortcut for ``SamplingPipeline(...).sample(logits)``. Shortcut for ``SamplingPipeline(...).sample(logits)``.
When **temperature** is exactly 0 (scalar or single-element tensor)
the function short-circuits to ``argmax`` for deterministic decode.
Args: Args:
logits: Raw logits ``[batch, vocab_size]``. logits: Raw logits ``[batch, vocab_size]``.
frequency_penalty: Penalty per occurrence for repeated tokens
(0.0 disables, range -2.0~2.0).
input_ids: Previously generated token IDs ``[batch, seq_len]``.
input_mask: Boolean mask for ``input_ids`` padding.
Returns: Returns:
Sampled token IDs ``[batch]``. Sampled token IDs ``[batch]``.
""" """
if SamplingPipeline._is_greedy(temperature):
return logits.argmax(dim=-1)
return SamplingPipeline( return SamplingPipeline(
[ [
TemperatureStrategy(temperature), TemperatureStrategy(temperature),
TopKStrategy(top_k), TopKStrategy(top_k),
TopPStrategy(top_p), TopPStrategy(top_p),
FrequencyPenaltyStrategy(frequency_penalty),
] ]
).sample(logits, filter_value) ).sample(logits, filter_value, input_ids, input_mask)
+9 -5
View File
@@ -6,7 +6,7 @@ import torch.nn.functional as F
from torch import Tensor from torch import Tensor
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.inference.core.cache import KvcacheView from astrai.inference.core.cache import CacheView
from astrai.model.components.linear import Linear from astrai.model.components.linear import Linear
from astrai.model.components.norm import RMSNorm from astrai.model.components.norm import RMSNorm
from astrai.model.components.rope import apply_rotary_emb from astrai.model.components.rope import apply_rotary_emb
@@ -38,6 +38,7 @@ class GQA(nn.Module):
norm_eps: float, norm_eps: float,
use_gated_attention: bool, use_gated_attention: bool,
layer_id: int, layer_id: int,
n_layers: int = 1,
): ):
super().__init__() super().__init__()
assert dim % n_heads == 0 assert dim % n_heads == 0
@@ -55,7 +56,7 @@ class GQA(nn.Module):
self.q_proj = Linear(dim, n_heads * self.head_dim) self.q_proj = Linear(dim, n_heads * self.head_dim)
self.k_proj = Linear(dim, n_kv_heads * self.head_dim) self.k_proj = Linear(dim, n_kv_heads * self.head_dim)
self.v_proj = Linear(dim, n_kv_heads * self.head_dim) self.v_proj = Linear(dim, n_kv_heads * self.head_dim)
self.o_proj = Linear(dim, dim) self.o_proj = Linear(dim, dim, init_std=0.02 / (2 * n_layers) ** 0.5)
if self.use_qk_norm: if self.use_qk_norm:
self.q_norm = RMSNorm(self.head_dim, norm_eps) self.q_norm = RMSNorm(self.head_dim, norm_eps)
@@ -74,7 +75,7 @@ class GQA(nn.Module):
x: Tensor, x: Tensor,
rotary_emb: Tensor, rotary_emb: Tensor,
attn_mask: Tensor = None, attn_mask: Tensor = None,
paged_cache: Optional[KvcacheView] = None, paged_cache: Optional[CacheView] = None,
) -> Tensor: ) -> Tensor:
is_causal = attn_mask is None is_causal = attn_mask is None
@@ -121,6 +122,7 @@ class MLA(nn.Module):
use_qk_norm: bool, use_qk_norm: bool,
use_gated_attention: bool, use_gated_attention: bool,
layer_id: int, layer_id: int,
n_layers: int = 1,
): ):
super().__init__() super().__init__()
self.dim = dim self.dim = dim
@@ -148,7 +150,9 @@ class MLA(nn.Module):
n_kv_heads * (2 * self.head_dim), n_kv_heads * (2 * self.head_dim),
) )
self.o_proj = Linear(dim, dim, bias=False) self.o_proj = Linear(
dim, dim, bias=False, init_std=0.02 / (2 * n_layers) ** 0.5
)
if use_gated_attention: if use_gated_attention:
self.gate = Linear(dim, dim, bias=False) self.gate = Linear(dim, dim, bias=False)
@@ -158,7 +162,7 @@ class MLA(nn.Module):
x: Tensor, x: Tensor,
rotary_emb: Tensor, rotary_emb: Tensor,
attn_mask: Tensor = None, attn_mask: Tensor = None,
paged_cache: Optional[KvcacheView] = None, paged_cache: Optional[CacheView] = None,
) -> Tensor: ) -> Tensor:
bsz, seq_len, _ = x.size() bsz, seq_len, _ = x.size()
is_causal = attn_mask is None is_causal = attn_mask is None
+10 -30
View File
@@ -1,51 +1,31 @@
from dataclasses import asdict
from typing import Optional from typing import Optional
import torch.nn as nn import torch.nn as nn
from torch import Tensor from torch import Tensor
from astrai.inference.core.cache import KvcacheView from astrai.inference.core.cache import CacheView
from astrai.model.components.attention import AttnFactory from astrai.model.components.attention import AttnFactory
from astrai.model.components.mlp import FFNFactory from astrai.model.components.mlp import FFNFactory
from astrai.model.components.norm import RMSNorm from astrai.model.components.norm import RMSNorm
class DecoderBlock(nn.Module): class DecoderBlock(nn.Module):
def __init__( def __init__(self, config, layer_id: int):
self,
dim: int,
n_heads: int,
dim_ffn: int,
n_kv_heads: int,
norm_eps: float,
use_qk_norm: bool,
use_gated_attention: bool,
layer_id: int,
attn_type: str = "gqa",
ffn_type: str = "mlp",
**kwargs,
):
super().__init__() super().__init__()
self.attention = AttnFactory.create( cfg = asdict(config)
attn_type, cfg["down_init_std"] = 0.02 / (2 * config.n_layers) ** 0.5
dim=dim, self.attention = AttnFactory.create(config.attn_type, **cfg, layer_id=layer_id)
n_heads=n_heads, self.input_norm = RMSNorm(config.dim, config.norm_eps)
n_kv_heads=n_kv_heads, self.post_attention_norm = RMSNorm(config.dim, config.norm_eps)
use_qk_norm=use_qk_norm, self.mlp = FFNFactory.create(config.ffn_type, **cfg)
norm_eps=norm_eps,
use_gated_attention=use_gated_attention,
layer_id=layer_id,
**kwargs,
)
self.input_norm = RMSNorm(dim, norm_eps)
self.post_attention_norm = RMSNorm(dim, norm_eps)
self.mlp = FFNFactory.create(ffn_type, dim, dim_ffn, **kwargs)
def forward( def forward(
self, self,
x: Tensor, x: Tensor,
rotary_emb: Tensor, rotary_emb: Tensor,
attention_mask: Optional[Tensor] = None, attention_mask: Optional[Tensor] = None,
paged_cache: Optional[KvcacheView] = None, paged_cache: Optional[CacheView] = None,
) -> Tensor: ) -> Tensor:
attn_output = self.attention( attn_output = self.attention(
self.input_norm(x), self.input_norm(x),
+5 -2
View File
@@ -7,10 +7,13 @@ from torch import Tensor
class Embedding(nn.Module): class Embedding(nn.Module):
def __init__(self, vocab_size: int, embedding_dim: int): def __init__(self, vocab_size: int, embedding_dim: int, neftune_alpha: float = 0.0):
super().__init__() super().__init__()
self.weight = nn.Parameter(torch.empty((vocab_size, embedding_dim))) self.weight = nn.Parameter(torch.empty((vocab_size, embedding_dim)))
self.neftune_noise_alpha = 0.0 self.neftune_noise_alpha = neftune_alpha
def set_neftune_alpha(self, alpha: float):
self.neftune_noise_alpha = alpha
def reset_parameters(self): def reset_parameters(self):
nn.init.normal_(self.weight, mean=0.0, std=0.02) nn.init.normal_(self.weight, mean=0.0, std=0.02)
+5 -2
View File
@@ -5,13 +5,16 @@ from torch import Tensor
class Linear(nn.Module): class Linear(nn.Module):
def __init__(self, in_dim: int, out_dim: int, bias: bool = False): def __init__(
self, in_dim: int, out_dim: int, bias: bool = False, init_std: float = 0.02
):
super().__init__() super().__init__()
self.weight = nn.Parameter(torch.empty((out_dim, in_dim))) self.weight = nn.Parameter(torch.empty((out_dim, in_dim)))
self.bias = nn.Parameter(torch.zeros(out_dim)) if bias else None self.bias = nn.Parameter(torch.zeros(out_dim)) if bias else None
self.init_std = init_std
def reset_parameters(self): def reset_parameters(self):
nn.init.kaiming_uniform_(self.weight, a=5**0.5) nn.init.normal_(self.weight, mean=0.0, std=self.init_std)
if self.bias is not None: if self.bias is not None:
fan_in, _ = nn.init._calculate_fan_in_and_fan_out(self.weight) fan_in, _ = nn.init._calculate_fan_in_and_fan_out(self.weight)
bound = 1 / (fan_in**0.5) bound = 1 / (fan_in**0.5)
+13 -4
View File
@@ -13,11 +13,11 @@ class FFNFactory(BaseFactory[nn.Module]):
@FFNFactory.register("mlp") @FFNFactory.register("mlp")
class MLP(nn.Module): class MLP(nn.Module):
def __init__(self, dim: int, dim_ffn: int): def __init__(self, dim: int, dim_ffn: int, down_init_std: float = 0.02):
super().__init__() super().__init__()
self.up = Linear(dim, dim_ffn) self.up = Linear(dim, dim_ffn)
self.gate = Linear(dim, dim_ffn) self.gate = Linear(dim, dim_ffn)
self.down = Linear(dim_ffn, dim) self.down = Linear(dim_ffn, dim, init_std=down_init_std)
def forward(self, x: Tensor) -> Tensor: def forward(self, x: Tensor) -> Tensor:
gated = self.up(x) * F.silu(self.gate(x)) gated = self.up(x) * F.silu(self.gate(x))
@@ -35,6 +35,7 @@ class DeepSeekMoE(nn.Module):
n_shared_experts: int = 1, n_shared_experts: int = 1,
n_activated_experts: int = 2, n_activated_experts: int = 2,
topk_method: str = "greedy", topk_method: str = "greedy",
n_layers: int = 1,
): ):
super().__init__() super().__init__()
self.dim = dim self.dim = dim
@@ -44,12 +45,20 @@ class DeepSeekMoE(nn.Module):
self.topk_method = topk_method self.topk_method = topk_method
self.router = Linear(dim, n_routed_experts, bias=False) self.router = Linear(dim, n_routed_experts, bias=False)
moe_scale = 1 / max(n_shared_experts, 1) + 1 / n_activated_experts
down_init_std = 0.02 / (2 * n_layers * moe_scale) ** 0.5
self.shared_experts = nn.ModuleList( self.shared_experts = nn.ModuleList(
[MLP(dim, dim_ffn) for _ in range(n_shared_experts)] [
MLP(dim, dim_ffn, down_init_std=down_init_std)
for _ in range(n_shared_experts)
]
) )
self.routed_experts = nn.ModuleList( self.routed_experts = nn.ModuleList(
[MLP(dim, dim_ffn) for _ in range(n_routed_experts)] [
MLP(dim, dim_ffn, down_init_std=down_init_std)
for _ in range(n_routed_experts)
]
) )
def forward(self, x: Tensor) -> Tensor: def forward(self, x: Tensor) -> Tensor:
+4 -14
View File
@@ -23,22 +23,12 @@ class EmbeddingEncoder(AutoModel):
self.rotary_embedding = RotaryEmbedding( self.rotary_embedding = RotaryEmbedding(
rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling
) )
self.embed_tokens = Embedding(config.vocab_size, config.dim) self.embed_tokens = Embedding(
config.vocab_size, config.dim, neftune_alpha=config.neftune_alpha
)
self.layers = nn.ModuleList( self.layers = nn.ModuleList(
[ [DecoderBlock(config, layer_id) for layer_id in range(config.n_layers)]
DecoderBlock(
config.dim,
config.n_heads,
config.dim_ffn,
config.n_kv_heads,
config.norm_eps,
config.use_qk_norm,
config.use_gated_attention,
layer_id,
)
for layer_id in range(config.n_layers)
]
) )
self.norm = RMSNorm(config.dim, config.norm_eps) self.norm = RMSNorm(config.dim, config.norm_eps)
+6 -25
View File
@@ -5,7 +5,7 @@ import torch.nn as nn
from torch import Tensor from torch import Tensor
from astrai.config.model_config import AutoRegressiveLMConfig from astrai.config.model_config import AutoRegressiveLMConfig
from astrai.inference.core.cache import KvcacheView from astrai.inference.core.cache import CacheView
from astrai.model.automodel import AutoModel from astrai.model.automodel import AutoModel
from astrai.model.components.decoder_block import DecoderBlock from astrai.model.components.decoder_block import DecoderBlock
from astrai.model.components.embedding import Embedding from astrai.model.components.embedding import Embedding
@@ -59,31 +59,12 @@ class AutoRegressiveLM(AutoModel):
self.rotary_embedding = RotaryEmbedding( self.rotary_embedding = RotaryEmbedding(
rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling
) )
self.embed_tokens = Embedding(config.vocab_size, config.dim) self.embed_tokens = Embedding(
config.vocab_size, config.dim, neftune_alpha=config.neftune_alpha
)
self.layers = nn.ModuleList( self.layers = nn.ModuleList(
[ [DecoderBlock(config, layer_id) for layer_id in range(config.n_layers)]
DecoderBlock(
config.dim,
config.n_heads,
config.dim_ffn,
config.n_kv_heads,
config.norm_eps,
config.use_qk_norm,
config.use_gated_attention,
layer_id,
attn_type=config.attn_type,
ffn_type=config.ffn_type,
n_routed_experts=config.n_routed_experts,
n_shared_experts=config.n_shared_experts,
n_activated_experts=config.n_activated_experts,
topk_method=config.topk_method,
kv_lora_rank=config.kv_lora_rank,
qk_nope_head_dim=config.qk_nope_head_dim,
qk_rope_head_dim=config.qk_rope_head_dim,
)
for layer_id in range(config.n_layers)
]
) )
self.norm = RMSNorm(config.dim, config.norm_eps) self.norm = RMSNorm(config.dim, config.norm_eps)
@@ -131,7 +112,7 @@ class AutoRegressiveLM(AutoModel):
self, self,
input_ids: Tensor, input_ids: Tensor,
input_mask: Optional[Tensor] = None, input_mask: Optional[Tensor] = None,
paged_cache: Optional[KvcacheView] = None, paged_cache: Optional[CacheView] = None,
position_ids: Optional[Tensor] = None, position_ids: Optional[Tensor] = None,
) -> Dict[str, Tensor]: ) -> Dict[str, Tensor]:
assert input_ids.ndim == 2 assert input_ids.ndim == 2
+40 -1
View File
@@ -7,6 +7,7 @@ from contextlib import contextmanager
from typing import Optional, Tuple from typing import Optional, Tuple
import torch import torch
import torch.distributed as dist
import torch.nn as nn import torch.nn as nn
from torch.distributed.fsdp import FullStateDictConfig, StateDictType from torch.distributed.fsdp import FullStateDictConfig, StateDictType
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
@@ -120,6 +121,21 @@ class BaseExecutor:
def unwrap_model(self, model: nn.Module): def unwrap_model(self, model: nn.Module):
return model.state_dict() return model.state_dict()
@contextmanager
def checkpoint_context(self, model: nn.Module):
if self.use_distributed:
dist.barrier()
state_dict = self._gather_state_dict(model)
yield state_dict
if self.use_distributed:
dist.barrier()
def _gather_state_dict(self, model: nn.Module):
state_dict = self.unwrap_model(model)
if self.use_distributed and get_rank() != 0:
return None
return state_dict
@property @property
def use_distributed(self) -> bool: def use_distributed(self) -> bool:
return get_world_size() > 1 return get_world_size() > 1
@@ -132,6 +148,19 @@ class BaseExecutor:
def grad_accum_steps(self) -> int: def grad_accum_steps(self) -> int:
return self.gradient_state.num_steps 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()
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
class ExecutorFactory(BaseFactory[BaseExecutor]): class ExecutorFactory(BaseFactory[BaseExecutor]):
pass pass
@@ -260,12 +289,22 @@ class FSDPExecutor(BaseExecutor):
return model.no_sync() return model.no_sync()
return contextlib.nullcontext() 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)
if isinstance(model, FSDP) and self.use_distributed:
total_norm = model.clip_grad_norm_(max_norm)
if isinstance(total_norm, torch.Tensor):
return total_norm.item()
return total_norm
return super().clip_grad_norm(model, max_norm)
def unwrap_model(self, model: nn.Module): def unwrap_model(self, model: nn.Module):
if isinstance(model, FSDP) and self.use_distributed: if isinstance(model, FSDP) and self.use_distributed:
with FSDP.state_dict_type( with FSDP.state_dict_type(
model, model,
StateDictType.FULL_STATE_DICT, StateDictType.FULL_STATE_DICT,
FullStateDictConfig(offload_to_cpu=True, rank0_only=False), FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
): ):
return model.state_dict() return model.state_dict()
+16 -5
View File
@@ -1,14 +1,21 @@
import os import os
import socket
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from contextlib import contextmanager from contextlib import contextmanager
from functools import wraps from functools import wraps
from typing import Callable from typing import Callable, Optional
import torch import torch
import torch.distributed as dist import torch.distributed as dist
import torch.multiprocessing as mp import torch.multiprocessing as mp
def find_free_port() -> str:
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
s.bind(("", 0))
return str(s.getsockname()[1])
def get_current_device(): def get_current_device():
return os.environ["LOCAL_DEVICE"] return os.environ["LOCAL_DEVICE"]
@@ -58,9 +65,11 @@ def setup_parallel(
os.environ["WORLD_SIZE"] = str(world_size) os.environ["WORLD_SIZE"] = str(world_size)
os.environ["LOCAL_DEVICE"] = str(device_id) os.environ["LOCAL_DEVICE"] = str(device_id)
dist.init_process_group( pg_kwargs = dict(rank=rank, world_size=world_size, backend=backend)
rank=rank, world_size=world_size, backend=backend, device_id=device_id if backend in ("nccl", "ccl"):
) pg_kwargs["device_id"] = device_id
dist.init_process_group(**pg_kwargs)
try: try:
if backend == "nccl" and torch.cuda.is_available(): if backend == "nccl" and torch.cuda.is_available():
@@ -215,11 +224,13 @@ def spawn_parallel_fn(
world_size: int, world_size: int,
backend: str = "nccl", backend: str = "nccl",
master_addr: str = "localhost", master_addr: str = "localhost",
master_port: str = "29500", master_port: Optional[str] = None,
device_type: str = "cuda", device_type: str = "cuda",
start_method: str = "spawn", start_method: str = "spawn",
**kwargs, **kwargs,
): ):
if master_port is None:
master_port = find_free_port()
launcher = _detect_launcher() launcher = _detect_launcher()
if launcher in ("torchelastic", "torchrun", "external"): if launcher in ("torchelastic", "torchrun", "external"):
strategy = TorchrunStrategy( strategy = TorchrunStrategy(
+8
View File
@@ -1,17 +1,21 @@
from astrai.preprocessing.builder import ( from astrai.preprocessing.builder import (
BaseMaskBuilder, BaseMaskBuilder,
MaskBuilderFactory, MaskBuilderFactory,
MultiOutputMaskBuilder,
SectionedMaskBuilder, SectionedMaskBuilder,
SingleOutputMaskBuilder,
) )
from astrai.preprocessing.packing import ( from astrai.preprocessing.packing import (
PackingStrategy, PackingStrategy,
PackingStrategyFactory, PackingStrategyFactory,
plan_bfd,
) )
from astrai.preprocessing.pipeline import Pipeline, filter_by_length from astrai.preprocessing.pipeline import Pipeline, filter_by_length
from astrai.preprocessing.position_id import ( from astrai.preprocessing.position_id import (
PositionIdStrategy, PositionIdStrategy,
PositionIdStrategyFactory, PositionIdStrategyFactory,
) )
from astrai.preprocessing.transform import TokenizeTransform
from astrai.preprocessing.writer import ( from astrai.preprocessing.writer import (
StoreWriter, StoreWriter,
StoreWriterFactory, StoreWriterFactory,
@@ -20,13 +24,17 @@ from astrai.preprocessing.writer import (
__all__ = [ __all__ = [
"BaseMaskBuilder", "BaseMaskBuilder",
"MaskBuilderFactory", "MaskBuilderFactory",
"MultiOutputMaskBuilder",
"PackingStrategy", "PackingStrategy",
"PackingStrategyFactory", "PackingStrategyFactory",
"Pipeline", "Pipeline",
"PositionIdStrategy", "PositionIdStrategy",
"PositionIdStrategyFactory", "PositionIdStrategyFactory",
"SectionedMaskBuilder", "SectionedMaskBuilder",
"SingleOutputMaskBuilder",
"StoreWriter", "StoreWriter",
"StoreWriterFactory", "StoreWriterFactory",
"TokenizeTransform",
"filter_by_length", "filter_by_length",
"plan_bfd",
] ]
+75 -53
View File
@@ -1,8 +1,10 @@
"""Mask building for preprocessing pipeline. """Mask building for preprocessing pipeline.
:class:`SectionRenderer` converts section specs into token ids and loss :class:`SectionRenderer` converts section specs into token ids and loss
masks (template / text / value extraction). :class:`SectionedMaskBuilder` masks (template / text / value extraction). :class:`SingleOutputMaskBuilder`
orchestrates single-output / multi-output (DPO / GRPO) assembly. handles single-output (SFT / pretrain), :class:`MultiOutputMaskBuilder`
handles multi-output (DPO / GRPO), and :class:`SectionedMaskBuilder`
orchestrates both modes as a façade.
""" """
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
@@ -93,8 +95,15 @@ class SectionRenderer:
return all_ids, loss_mask return all_ids, loss_mask
def process_list_field(self, item: dict, sections: list, config, tokenizer): def process_list_field(self, item: dict, sections: list, config, tokenizer):
all_ids: list[int] = [] """Tokenize a list-valued field, preserving per-element boundaries.
loss_mask: list[int] = []
Returns ``(list_of_id_lists, list_of_mask_lists)`` where each
inner list corresponds to one element of the source list. This
is critical for GRPO where each response must stay a separate
sequence so the strategy can form a ``[G, R]`` tensor.
"""
per_item_ids: list[list[int]] = []
per_item_masks: list[list[int]] = []
for sec in sections: for sec in sections:
field = sec["field"] field = sec["field"]
@@ -106,17 +115,13 @@ class SectionRenderer:
continue continue
for val in values: for val in values:
ids: list[int] = []
mask: list[int] = []
if use_template: if use_template:
if isinstance(val, list): if isinstance(val, list):
wrapper = {field: val} wrapper = {field: val}
self._append_template( self._append_template(
wrapper, wrapper, field, action, tokenizer, config, ids, mask
field,
action,
tokenizer,
config,
all_ids,
loss_mask,
) )
else: else:
wrapper = {field: str(val)} wrapper = {field: str(val)}
@@ -128,17 +133,19 @@ class SectionRenderer:
False, False,
False, False,
config, config,
all_ids, ids,
loss_mask, mask,
) )
if ids:
max_len = config.preprocessing.max_seq_len max_len = config.preprocessing.max_seq_len
all_ids = all_ids[:max_len] ids = ids[:max_len]
loss_mask = loss_mask[: len(all_ids)] mask = mask[: len(ids)]
per_item_ids.append(ids)
per_item_masks.append(mask)
if not all_ids: if not per_item_ids:
return None, None return None, None
return all_ids, loss_mask return per_item_ids, per_item_masks
@staticmethod @staticmethod
def is_value_section(sections: list) -> bool: def is_value_section(sections: list) -> bool:
@@ -212,42 +219,17 @@ class MaskBuilderFactory(BaseFactory["BaseMaskBuilder"]):
pass pass
@MaskBuilderFactory.register("sectioned") @MaskBuilderFactory.register("single")
class SectionedMaskBuilder(BaseMaskBuilder): class SingleOutputMaskBuilder(BaseMaskBuilder):
"""Config-driven builder supporting single and multi-output modes. """Build a single output sequence with optional loss mask.
Single-output:: Expects ``config.input.sections`` (list of section specs).
{"input": {"sections": [
{"field": "messages", "action": "$role", "template": true}
]}}
{"sequence": [...], "loss_mask": [...], "domain": "..."}
Multi-output (DPO / GRPO)::
{"input": {"sources": {
"chosen": {"sections": [{"field": "chosen", "action": "$role", "template": true}]},
"rejected": {"sections": [{"field": "rejected", "action": "$role", "template": true}]},
}}}
{"chosen": [...], "chosen_mask": [...], "rejected": [...], "rejected_mask": [...], "domain": "..."}
Output spec fields::
sections list of section specs (same format as single-output)
list_field True when JSONL field holds a list (GRPO responses)
mask_key explicit loss-mask output key (default: ``"{output_key}_mask"``)
""" """
def __init__(self): def __init__(self, renderer: Optional[SectionRenderer] = None):
self.renderer = SectionRenderer() self.renderer = renderer or SectionRenderer()
def build(self, item: dict, config, tokenizer) -> Optional[dict]: def build(self, item: dict, config, tokenizer) -> Optional[dict]:
sources_spec = getattr(config.input, "sources", None)
if sources_spec:
return self._build_multi(item, sources_spec, config, tokenizer)
return self._build_single(item, config, tokenizer)
def _build_single(self, item: dict, config, tokenizer) -> Optional[dict]:
sections = config.input.sections sections = config.input.sections
if not sections: if not sections:
return None return None
@@ -266,9 +248,22 @@ class SectionedMaskBuilder(BaseMaskBuilder):
result["loss_mask"] = mask result["loss_mask"] = mask
return result return result
def _build_multi(
self, item: dict, sources_spec: dict, config, tokenizer @MaskBuilderFactory.register("multi")
) -> Optional[dict]: class MultiOutputMaskBuilder(BaseMaskBuilder):
"""Build multiple output sequences (DPO / GRPO).
Expects ``config.input.sources`` (dict of output_key → spec).
"""
def __init__(self, renderer: Optional[SectionRenderer] = None):
self.renderer = renderer or SectionRenderer()
def build(self, item: dict, config, tokenizer) -> Optional[dict]:
sources_spec = getattr(config.input, "sources", None)
if not sources_spec:
return None
result: dict = {} result: dict = {}
any_output = False any_output = False
@@ -292,7 +287,15 @@ class SectionedMaskBuilder(BaseMaskBuilder):
ids, mask = self.renderer.process_list_field( ids, mask = self.renderer.process_list_field(
item, sections, config, tokenizer item, sections, config, tokenizer
) )
else: if ids is None:
continue
# ids is List[List[int]] — preserve per-response structure
result[output_key] = ids
if mask is not None:
result[mask_key] = mask
any_output = True
continue
ids, mask = self.renderer.process_sections( ids, mask = self.renderer.process_sections(
item, sections, config, tokenizer, is_top_level=True item, sections, config, tokenizer, is_top_level=True
) )
@@ -313,3 +316,22 @@ class SectionedMaskBuilder(BaseMaskBuilder):
result["domain"] = _extract_domain(item, config.output.domain_key) result["domain"] = _extract_domain(item, config.output.domain_key)
return result return result
@MaskBuilderFactory.register("sectioned")
class SectionedMaskBuilder(BaseMaskBuilder):
"""Façade that dispatches to SingleOutputMaskBuilder or MultiOutputMaskBuilder.
Preserves backward compatibility for existing configs and code that rely
on the ``"sectioned"`` factory name.
"""
def __init__(self):
self._single = SingleOutputMaskBuilder()
self._multi = MultiOutputMaskBuilder()
def build(self, item: dict, config, tokenizer) -> Optional[dict]:
sources_spec = getattr(config.input, "sources", None)
if sources_spec:
return self._multi.build(item, config, tokenizer)
return self._single.build(item, config, tokenizer)
+124
View File
@@ -0,0 +1,124 @@
"""Shared preprocessing kernel used by both :class:`Pipeline` and
:class:`TokenizeTransform`.
The two entry points previously duplicated ~60 % of their logic:
record iteration, mask-builder invocation, primary-id extraction,
per-key accumulation, dtype inference and position-id generation.
This module factors out the common core as pure functions so that
the online (``TokenizeTransform``) and offline (``Pipeline``) paths
stay in lockstep.
"""
from itertools import chain
from typing import Dict, Iterator, List, Optional
import torch
from astrai.config.preprocess_config import PipelineConfig
from astrai.preprocessing.builder import MaskBuilderFactory
from astrai.preprocessing.position_id import PositionIdStrategyFactory
from astrai.tokenize import AutoTokenizer
def build_preprocessing_components(config: PipelineConfig, tokenizer_path: str):
"""Load tokenizer, mask builder and position-id strategy together.
Both ``Pipeline`` and ``TokenizeTransform`` need the same triple;
centralising the construction avoids drift (e.g. one path forgetting
to create the position-id strategy).
"""
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
mask_builder = MaskBuilderFactory.create("sectioned")
position_strategy = PositionIdStrategyFactory.create(
config.output.position_ids_mode
)
return tokenizer, mask_builder, position_strategy
def primary_ids(result: dict) -> List[int]:
"""Return the first flat int-list value in *result*.
Used for token counting and position-id generation when the
primary key name is not known (DPO uses ``chosen``, GRPO uses
``prompts``, SFT uses ``sequence``).
"""
for val in result.values():
if isinstance(val, list) and val and isinstance(val[0], int):
return val
return []
def infer_dtype(ids: List) -> torch.dtype:
"""Float values become float32, everything else int32."""
if ids and isinstance(ids[0], float):
return torch.float32
return torch.int32
def iter_raw_records(
records: List[dict],
mask_builder,
config: PipelineConfig,
tokenizer,
) -> Iterator[dict]:
"""Yield mask-builder output dicts for each record, skipping failures.
Drops ``domain`` from the result (callers that need it should read
it before calling this). Each yielded dict maps a key
(``sequence``, ``chosen``, ``responses``…) to either a flat
``List[int]`` or a nested ``List[List[int]]`` (GRPO responses/masks).
"""
for item in records:
result = mask_builder.build(item, config, tokenizer)
if result is None:
continue
result.pop("domain", None)
if not primary_ids(result):
continue
yield result
def to_per_record_tensors(
raw: Dict[str, list],
) -> Dict[str, List[torch.Tensor]]:
"""Convert an accumulated ``{key: [per-record ids]}`` dict to tensors.
Handles three shapes transparently:
- ``List[int]`` per record (``sequence``, ``chosen``…) → one tensor per record.
- ``List[List[int]]`` per record (GRPO ``responses``/``masks``) → one
``List[Tensor]`` per record (nested), preserving the per-response
boundary so downstream code can index responses individually.
- ``List[int]`` for the whole shard (pre-packed keys) → single tensor.
The detection mirrors the previous inline logic in
``Pipeline._flush`` and ``TokenizeTransform.apply``.
"""
tensors: Dict[str, List[torch.Tensor]] = {}
for key, ids_list in raw.items():
if ids_list and isinstance(ids_list[0], list):
tensors[key] = [
[torch.tensor(sub, dtype=infer_dtype(sub)) for sub in ids]
if ids and isinstance(ids[0], list)
else torch.tensor(ids, dtype=infer_dtype(ids))
for ids in ids_list
]
else:
tensors[key] = [
torch.tensor(list(chain.from_iterable(ids_list)), dtype=torch.int32)
]
return tensors
def build_position_ids(
sequences: List[List[int]],
strategy,
) -> Optional[List[int]]:
"""Generate position ids for *sequences* using *strategy*.
Returns ``None`` when the strategy produces no ids (e.g. ``none``
mode), so callers can skip attaching the key instead of storing
an empty list.
"""
pos_ids = strategy.generate(sequences)
return pos_ids or None
+83 -28
View File
@@ -19,6 +19,43 @@ def _truncate(seq: List[int], max_len: int, mode: str) -> List[int]:
return seq[:max_len] return seq[:max_len]
def plan_bfd(
sequences: List[List[int]], max_packed_len: int, truncation_mode: str = "keep_start"
) -> List[List[int]]:
"""Best-Fit Decreasing bin packing of *sequences* into bins.
Returns a list of bins, each bin a list of original indices into
*sequences*. Bin capacities are respected on the *truncated*
length of each sequence (so a sequence longer than
*max_packed_len* counts at *max_packed_len*).
Pure index-based so callers can apply the same plan to any
aligned key (``loss_mask``, ``position_ids``…).
"""
n = len(sequences)
order = sorted(range(n), key=lambda i: len(sequences[i]), reverse=True)
bins: List[List[int]] = []
bin_lengths: List[int] = []
for orig_idx in order:
seq_len = len(_truncate(sequences[orig_idx], max_packed_len, truncation_mode))
best_bin = None
best_remain = max_packed_len + 1
for i, bl in enumerate(bin_lengths):
remain = max_packed_len - bl
if seq_len <= remain < best_remain:
best_remain = remain
best_bin = i
if best_bin is not None:
bins[best_bin].append(orig_idx)
bin_lengths[best_bin] += seq_len
else:
bins.append([orig_idx])
bin_lengths.append(seq_len)
return bins
class PackingStrategy(ABC): class PackingStrategy(ABC):
"""Reorder and truncate sequences within a shard.""" """Reorder and truncate sequences within a shard."""
@@ -70,7 +107,7 @@ class BFDPacking(PackingStrategy):
sequences = keys.get("sequence", []) sequences = keys.get("sequence", [])
if not sequences: if not sequences:
return keys return keys
bins = self._plan(sequences, max_packed_len, truncation_mode) bins = plan_bfd(sequences, max_packed_len, truncation_mode)
packed: Dict[str, List[List[int]]] = {} packed: Dict[str, List[List[int]]] = {}
for k, vals in keys.items(): for k, vals in keys.items():
@@ -91,31 +128,49 @@ class BFDPacking(PackingStrategy):
result.extend(vals[i]) result.extend(vals[i])
return result return result
@PackingStrategyFactory.register("bfd_split")
class BFDSplitPacking(BFDPacking):
"""BFD packing with over-length sequences split into chunks.
Sequences longer than *max_packed_len* are split into consecutive
chunks of at most *max_packed_len* tokens instead of being
truncated. Each chunk becomes an independent sequence that enters
BFD planning. All keys (``loss_mask``, ``position_ids``, …) are
split in lockstep so per-token alignment is preserved.
Note: because each chunk is treated as a separate document, the
second chunk of a split sequence loses the preceding context.
"""
def apply(
self,
keys: Dict[str, List[List[int]]],
max_packed_len: int,
truncation_mode: str,
) -> Dict[str, List[List[int]]]:
sequences = keys.get("sequence", [])
if not sequences:
return keys
if max_packed_len <= 0:
return super().apply(keys, max_packed_len, truncation_mode)
split_keys = self._split_all(keys, max_packed_len)
return super().apply(split_keys, max_packed_len, truncation_mode)
@staticmethod @staticmethod
def _plan( def _split_all(
sequences: List[List[int]], max_packed_len: int, truncation_mode: str keys: Dict[str, List[List[int]]], max_packed_len: int
) -> List[List[int]]: ) -> Dict[str, List[List[int]]]:
n = len(sequences) """Split every sequence exceeding *max_packed_len* into chunks,
order = sorted(range(n), key=lambda i: len(sequences[i]), reverse=True) applying the same chunk boundaries to all keys."""
bins: List[List[int]] = [] sequences = keys["sequence"]
bin_lengths: List[int] = [] chunk_bounds = [list(range(0, len(s), max_packed_len)) for s in sequences]
result: Dict[str, List[List[int]]] = {}
for orig_idx in order: for key, vals in keys.items():
seq_len = len( split_vals: List[List[int]] = []
_truncate(sequences[orig_idx], max_packed_len, truncation_mode) for val, starts in zip(vals, chunk_bounds):
) for start in starts:
best_bin = None split_vals.append(val[start : start + max_packed_len])
best_remain = max_packed_len + 1 result[key] = split_vals
for i, bl in enumerate(bin_lengths): return result
remain = max_packed_len - bl
if seq_len <= remain < best_remain:
best_remain = remain
best_bin = i
if best_bin is not None:
bins[best_bin].append(orig_idx)
bin_lengths[best_bin] += seq_len
else:
bins.append([orig_idx])
bin_lengths.append(seq_len)
return bins
+102 -39
View File
@@ -4,6 +4,10 @@ Composes a :class:`BaseMaskBuilder` (selected by ``input.type``) with
sharding and flush to ``.h5`` / ``.bin`` storage. Packing, position-id sharding and flush to ``.h5`` / ``.bin`` storage. Packing, position-id
generation and storage writing are each delegated to pluggable strategies, generation and storage writing are each delegated to pluggable strategies,
dispatched by configuration keys. dispatched by configuration keys.
Record iteration, mask building, primary-id extraction and per-key
accumulation are shared with :class:`TokenizeTransform` via the
:mod:`astrai.preprocessing.core` helpers.
""" """
import json import json
@@ -17,11 +21,13 @@ import torch
import tqdm import tqdm
from astrai.config.preprocess_config import PipelineConfig from astrai.config.preprocess_config import PipelineConfig
from astrai.preprocessing.builder import MaskBuilderFactory from astrai.preprocessing.core import (
build_preprocessing_components,
iter_raw_records,
primary_ids,
)
from astrai.preprocessing.packing import PackingStrategyFactory from astrai.preprocessing.packing import PackingStrategyFactory
from astrai.preprocessing.position_id import PositionIdStrategyFactory
from astrai.preprocessing.writer import StoreWriterFactory from astrai.preprocessing.writer import StoreWriterFactory
from astrai.tokenize import AutoTokenizer
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -64,20 +70,18 @@ class Pipeline:
self.output_dir = output_dir self.output_dir = output_dir
self.tokenizer_path = tokenizer_path self.tokenizer_path = tokenizer_path
self.mask_builder = MaskBuilderFactory.create("sectioned") self.tokenizer, self.mask_builder, self._position_id = (
build_preprocessing_components(config, tokenizer_path)
)
self._packer = PackingStrategyFactory.create( self._packer = PackingStrategyFactory.create(
config.preprocessing.packing_strategy config.preprocessing.packing_strategy
) )
self._position_id = PositionIdStrategyFactory.create(
config.output.position_ids_mode
)
self._writer = StoreWriterFactory.create(config.output.storage_format) self._writer = StoreWriterFactory.create(config.output.storage_format)
def transform(self, item: dict) -> Optional[dict]: def transform(self, item: dict) -> Optional[dict]:
return self.mask_builder.build(item, self.config, self._tokenizer) return self.mask_builder.build(item, self.config, self.tokenizer)
def run(self): def run(self):
self._tokenizer = AutoTokenizer.from_pretrained(self.tokenizer_path)
domains: dict = defaultdict(lambda: defaultdict(list)) domains: dict = defaultdict(lambda: defaultdict(list))
total_tokens = 0 total_tokens = 0
shard_idx: dict[str, int] = defaultdict(int) shard_idx: dict[str, int] = defaultdict(int)
@@ -102,14 +106,7 @@ class Pipeline:
continue continue
domain = result.pop("domain", "__default__") domain = result.pop("domain", "__default__")
ids = primary_ids(result)
is_multi = bool(getattr(self.config.input, "sources", None))
if is_multi:
ids = self._primary_ids(result)
else:
ids = result.pop("sequence")
result["sequence"] = ids
if not ids: if not ids:
continue continue
@@ -129,26 +126,24 @@ class Pipeline:
if total_tokens > 0: if total_tokens > 0:
self._flush(domains, shard_idx) self._flush(domains, shard_idx)
@staticmethod
def _primary_ids(result: dict) -> list:
"""Return the first list-valued entry in *result* as the primary id
sequence for token counting."""
for val in result.values():
if isinstance(val, list) and val and isinstance(val[0], int):
return val
return []
@staticmethod @staticmethod
def _align_bucket(bucket: dict, result: dict, ids: list): def _align_bucket(bucket: dict, result: dict, ids: list):
"""Pad previously-accumulated keys that are missing from *result*.""" """Pad previously-accumulated keys that are missing from *result*."""
for key in list(bucket.keys()): for key in list(bucket.keys()):
if key in result: if key in result:
continue continue
bucket[key].append([1] * len(ids)) bucket[key].append([0] * len(ids))
def _iter_items(self): def _iter_items(self):
for path in self.paths: for path in self.paths:
with open(path, "r", encoding="utf-8") as f: with open(path, "r", encoding="utf-8") as f:
if path.endswith(".json"):
data = json.load(f)
if isinstance(data, dict):
yield data
elif isinstance(data, list):
yield from data
else:
for line in f: for line in f:
line = line.strip() line = line.strip()
if not line: if not line:
@@ -160,20 +155,15 @@ class Pipeline:
idx = shard_idx[domain] idx = shard_idx[domain]
pp = self.config.preprocessing pp = self.config.preprocessing
original_sequences = keys.get("sequence", [])
mode = self.config.output.position_ids_mode
keys = self._inject_doc_reset_position_ids(keys, mode, original_sequences)
keys = self._packer.apply(dict(keys), pp.max_packed_len, pp.truncation_mode) keys = self._packer.apply(dict(keys), pp.max_packed_len, pp.truncation_mode)
tensors = self._to_tensors(keys)
tensors: Dict[str, List[torch.Tensor]] = {} tensors = self._inject_continuous_position_ids(
for key, ids_list in keys.items(): tensors, mode, keys.get("sequence", [])
dt = _STR_TO_DTYPE.get(
self.config.output.dtype.get(key, "int32"), torch.int32
) )
tensors[key] = [
torch.tensor(list(chain.from_iterable(ids_list)), dtype=dt)
]
pos_ids = self._position_id.generate(keys.get("sequence", []))
if pos_ids:
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
self._writer.save(self.output_dir, domain, idx, tensors) self._writer.save(self.output_dir, domain, idx, tensors)
shard_idx[domain] = idx + 1 shard_idx[domain] = idx + 1
@@ -183,3 +173,76 @@ class Pipeline:
f" saved {domain}/shard_{idx:04d} " f" saved {domain}/shard_{idx:04d} "
f"({tensors[first_key][0].numel():,} tokens)" f"({tensors[first_key][0].numel():,} tokens)"
) )
def _inject_doc_reset_position_ids(
self,
keys: Dict[str, list],
mode: str,
original_sequences: List[List[int]],
) -> Dict[str, list]:
"""Attach per-document position_ids before packing (``doc_reset``).
``doc_reset`` position ids must enter the packer so that each
packed bin concatenates the per-doc ranges in bin order. The
per-record structure ``[range(len(s)) for s in seqs]`` is required
by the packer (it concatenates per-record lists per bin); the
``PositionIdStrategy.generate`` flattens, so it cannot be used
directly here it is only consulted for the ``continuous``
post-packing path.
"""
if mode != "doc_reset" or not original_sequences:
return keys
keys["position_ids"] = [list(range(len(s))) for s in original_sequences]
return keys
def _inject_continuous_position_ids(
self,
tensors: Dict[str, List[torch.Tensor]],
mode: str,
packed_sequences: List[List[int]],
) -> Dict[str, List[torch.Tensor]]:
"""Attach a single continuous position_ids tensor after packing.
``continuous`` mode spans the whole shard (post-packing), so it
cannot participate in bin packing it is computed from the
packed sequences and appended directly to the tensor dict.
"""
if mode != "continuous" or not packed_sequences:
return tensors
pos_ids = self._position_id.generate(packed_sequences)
if pos_ids:
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
return tensors
def _to_tensors(self, keys: Dict[str, list]) -> Dict[str, List[torch.Tensor]]:
"""Convert packed per-key id lists to tensors.
Honours ``config.output.dtype`` overrides per key; falls back to
``int32``. Handles three shapes (see
:func:`astrai.preprocessing.core.to_per_record_tensors` for the
equivalent online-path helper):
- ``List[int]`` per record one tensor per record.
- ``List[List[int]]`` per record (GRPO responses/masks) one tensor
per record, inner lists flattened.
- ``List[int]`` for the whole shard (pre-packed keys) single tensor.
"""
tensors: Dict[str, List[torch.Tensor]] = {}
for key, ids_list in keys.items():
dt = _STR_TO_DTYPE.get(
self.config.output.dtype.get(key, "int32"), torch.int32
)
if ids_list and isinstance(ids_list[0], list):
tensors[key] = [
torch.tensor(
list(chain.from_iterable(ids))
if ids and isinstance(ids[0], list)
else ids,
dtype=dt,
)
for ids in ids_list
]
else:
tensors[key] = [
torch.tensor(list(chain.from_iterable(ids_list)), dtype=dt)
]
return tensors
+92
View File
@@ -0,0 +1,92 @@
"""Tokenization transform for JSONL record streams.
Bridges the Reader layer (``JsonlStore`` reads raw JSON records) and the
Dataset layer (expects per-record tensors). Holds the tokenizer,
mask-builder and position-id strategy together so that I/O code stays
free of model dependencies.
The record-processing core (mask building, primary-id extraction,
per-key tensorisation, position-id generation) is shared with
:class:`astrai.preprocessing.pipeline.Pipeline` via the
:mod:`astrai.preprocessing.core` helpers.
"""
import json
from pathlib import Path
from typing import Dict, List
import torch
from astrai.config.preprocess_config import PipelineConfig
from astrai.preprocessing.core import (
build_position_ids,
build_preprocessing_components,
iter_raw_records,
to_per_record_tensors,
)
class TokenizeTransform:
"""Tokenize raw JSONL record dicts into per-key tensor lists.
Owns the three preprocessing concerns that were previously inlined in
``JsonlStore``: tokenization, loss-mask construction and position-id
generation. Constructing it loads the tokenizer, so it is intentionally
cheap to pass around once built.
Args:
config: Pipeline config describing sections / masks / position mode.
tokenizer_path: Path passed to ``AutoTokenizer.from_pretrained``.
"""
def __init__(self, config: PipelineConfig, tokenizer_path: str):
self.config = config
self.tokenizer, self.mask_builder, self.position_strategy = (
build_preprocessing_components(config, tokenizer_path)
)
@classmethod
def from_config_file(cls, config_path: str) -> "TokenizeTransform":
"""Build from a ``dataset_config.json`` file path.
The config file follows :class:`PipelineConfig` schema with an
extra ``tokenizer_path`` field. When omitted, the config's
parent directory is used as the tokenizer path.
"""
root = Path(config_path).parent
with open(config_path, "r", encoding="utf-8") as f:
raw_config = json.load(f)
tokenizer_path = raw_config.pop("tokenizer_path", None) or str(root)
config = PipelineConfig.from_dict(raw_config)
return cls(config, tokenizer_path)
def apply(self, records: List[dict]) -> Dict[str, list]:
"""Tokenize a list of raw record dicts.
Returns a dict mapping key (``sequence``, ``chosen``, ``responses``,
) to a list of per-record tensors (or nested tensor lists for
multi-response keys such as GRPO ``responses``).
"""
raw: Dict[str, list] = {}
doc_sequences: List[List[int]] = []
for result in iter_raw_records(
records, self.mask_builder, self.config, self.tokenizer
):
primary = None
for val in result.values():
if isinstance(val, list) and val and isinstance(val[0], int):
primary = val
break
if primary is not None:
doc_sequences.append(primary)
for key, ids in result.items():
raw.setdefault(key, []).append(ids)
tensors = to_per_record_tensors(raw)
pos_ids = build_position_ids(doc_sequences, self.position_strategy)
if pos_ids is not None:
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
return tensors
+1 -1
View File
@@ -14,8 +14,8 @@ from typing import Dict, List
import torch import torch
from astrai.dataset.storage import save_bin, save_h5
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.serialization import save_bin, save_h5
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
+45
View File
@@ -0,0 +1,45 @@
"""Serialization utilities for models and datasets.
This package re-exports checkpoint helpers and dataset storage helpers so
that existing imports from ``astrai.serialization`` continue to work.
"""
from astrai.serialization.checkpoint import (
Checkpoint,
load_json,
load_model_config,
load_model_weights,
load_safetensors,
load_state_dict,
load_torch,
save_json,
save_model,
save_safetensors,
save_torch,
)
from astrai.serialization.dataset import (
load_bin,
load_bin_offsets,
load_h5,
save_bin,
save_h5,
)
__all__ = [
"Checkpoint",
"load_json",
"load_model_config",
"load_model_weights",
"load_safetensors",
"load_state_dict",
"load_torch",
"save_json",
"save_model",
"save_safetensors",
"save_torch",
"load_bin",
"load_bin_offsets",
"load_h5",
"save_bin",
"save_h5",
]
@@ -1,3 +1,5 @@
"""Model checkpoint serialization helpers."""
import io import io
import json import json
import time import time
@@ -136,7 +138,7 @@ def load_state_dict(path: Union[str, Path], broadcast: bool = False) -> dict:
class Checkpoint: class Checkpoint:
state_dict: Dict[str, Any] = field(default_factory=dict) state_dict: Dict[str, Any] = field(default_factory=dict)
epoch: int = 0 epoch: int = 0
iteration: int = 0 consumed_samples: int = 0
extra: Dict[str, Any] = field(default_factory=dict) extra: Dict[str, Any] = field(default_factory=dict)
meta: Dict[str, Any] = field(default_factory=dict) meta: Dict[str, Any] = field(default_factory=dict)
config: Dict[str, Any] = field(default_factory=dict) config: Dict[str, Any] = field(default_factory=dict)
@@ -145,12 +147,9 @@ class Checkpoint:
save_path = Path(save_dir) save_path = Path(save_dir)
save_path.mkdir(parents=True, exist_ok=True) save_path.mkdir(parents=True, exist_ok=True)
if get_rank() != 0:
return
meta = { meta = {
"epoch": self.epoch, "epoch": self.epoch,
"iteration": self.iteration, "consumed_samples": self.consumed_samples,
"timestamp": time.strftime("%Y-%m-%dT%H:%M:%S"), "timestamp": time.strftime("%Y-%m-%dT%H:%M:%S"),
**self.meta, **self.meta,
} }
@@ -176,8 +175,9 @@ class Checkpoint:
return cls( return cls(
state_dict=state_dict, state_dict=state_dict,
epoch=meta.get("epoch", 0), epoch=meta.get("epoch", 0),
iteration=meta.get("iteration", 0), consumed_samples=meta.get("consumed_samples", 0),
extra=extra, extra=extra,
meta=meta,
config=config, config=config,
) )
+123
View File
@@ -0,0 +1,123 @@
"""Dataset storage serialization helpers (HDF5 / memory-mapped binary)."""
import json
import os
from pathlib import Path
from typing import Any, Dict, List, Optional
import h5py
import numpy as np
import torch
from torch import Tensor
def save_h5(file_path: str, file_name: str, tensor_group: Dict[str, List[Tensor]]):
os.makedirs(file_path, exist_ok=True)
full_file_path = os.path.join(file_path, f"{file_name}.h5")
with h5py.File(full_file_path, "w") as f:
for key, tensors in tensor_group.items():
grp = f.create_group(key)
for idx, tensor in enumerate(tensors):
arr = tensor.cpu().numpy()
grp.create_dataset(f"data_{idx}", data=arr)
def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
tensor_group: Dict[str, List[Tensor]] = {}
root_path = Path(file_path)
if root_path.is_file() and root_path.suffix in (".h5", ".hdf5"):
h5_files = [root_path]
else:
h5_files = list(root_path.rglob("*.h5")) + list(root_path.rglob("*.hdf5"))
for h5_file in h5_files:
with h5py.File(h5_file, "r") as f:
for key in f.keys():
grp = f[key]
dsets = []
for dset_name in grp.keys():
dset = grp[dset_name]
tensor = torch.from_numpy(dset[:])
if share_memory:
tensor = tensor.share_memory_()
dsets.append(tensor)
if tensor_group.get(key) is None:
tensor_group[key] = []
tensor_group[key].extend(dsets)
return tensor_group
def save_bin(
file_path: str,
tensor_group: Dict[str, List[Tensor]],
record_keys: Optional[List[str]] = None,
):
"""Save tensors as memory-mapped binary files.
When *record_keys* is provided, those keys are written with per-record
cumulative offsets in ``meta.json`` so that ``MmapStore.fetch_record``
can slice individual records from the concatenated binary without
cross-record concatenation. Keys not in *record_keys* (e.g. SEQ
``sequence``) are written as a single contiguous stream without
offsets, preserving backward compatibility.
Nested keys (``List[List[Tensor]]`` such as GRPO ``responses``) are
not supported in bin format use H5 for those.
"""
os.makedirs(file_path, exist_ok=True)
record_keys = set(record_keys or [])
meta = {}
for key, tensors in tensor_group.items():
if tensors and isinstance(tensors[0], list):
raise ValueError(
f"Nested key '{key}' (List[List[Tensor]]) is not supported "
f"in bin format. Use H5 or JSONL storage instead."
)
cat = torch.cat(tensors, dim=0)
entry: Dict[str, Any] = {
"shape": list(cat.shape),
"dtype": str(cat.dtype).split(".")[-1],
}
if key in record_keys:
offsets = [0]
for t in tensors:
offsets.append(offsets[-1] + t.shape[0])
entry["offsets"] = offsets
meta[key] = entry
np.asarray(cat.cpu().numpy()).tofile(os.path.join(file_path, f"{key}.bin"))
with open(os.path.join(file_path, "meta.json"), "w") as f:
json.dump(meta, f)
def load_bin(file_path: str) -> Dict[str, List[Tensor]]:
with open(os.path.join(file_path, "meta.json"), "r") as f:
meta = json.load(f)
segments: Dict[str, List[Tensor]] = {}
for key, info in meta.items():
arr = np.memmap(
os.path.join(file_path, f"{key}.bin"),
dtype=info["dtype"],
mode="r",
shape=tuple(info["shape"]),
)
segments[key] = [torch.from_numpy(arr)]
return segments
def load_bin_offsets(file_path: str) -> Dict[str, List[int]]:
"""Read per-record cumulative offsets from ``meta.json``.
Returns an empty dict when no key has offsets (legacy bin files),
in which case record-mode access falls back to per-record segment
indexing (H5/JSONL layout).
"""
with open(os.path.join(file_path, "meta.json"), "r") as f:
meta = json.load(f)
offsets: Dict[str, List[int]] = {}
for key, info in meta.items():
if "offsets" in info:
offsets[key] = info["offsets"]
return offsets
+14 -1
View File
@@ -1,3 +1,4 @@
from functools import cached_property
from typing import Any, Dict, List, Optional from typing import Any, Dict, List, Optional
from jinja2 import Template from jinja2 import Template
@@ -29,7 +30,19 @@ class ChatTemplate:
self.description = description self.description = description
self.default_variables = default_variables or {} self.default_variables = default_variables or {}
self.special_tokens = special_tokens or {} self.special_tokens = special_tokens or {}
self._compiled: Template = Template(template_str)
@cached_property
def _compiled(self) -> Template:
"""Lazy-compiled Jinja2 template, cached on first access.
The compiled :class:`~jinja2.Template` holds a dynamically-generated
``root`` render function whose ``__module__`` is ``None``; under
``pickle`` it falls back to ``__main__`` and breaks ``spawn``-based
multiprocessing. By deferring compilation to first access, the
default pickle protocol serialises only ``template_str``; each
worker rebuilds the cache on first render.
"""
return Template(self.template_str)
@classmethod @classmethod
def from_string( def from_string(
+7
View File
@@ -164,7 +164,14 @@ class AutoTokenizer:
- tokenizer.bos_token returns string - tokenizer.bos_token returns string
- tokenizer.bos_token_id returns corresponding integer ID - tokenizer.bos_token_id returns corresponding integer ID
- tokenizer.stop_ids returns list of corresponding integer IDs for all special tokens - tokenizer.stop_ids returns list of corresponding integer IDs for all special tokens
Internal/private attrs are not intercepted: during unpickle
``__dict__`` is empty, so probing ``self._special_token_map``
would recurse infinitely.
""" """
if key.startswith("_"):
raise AttributeError(key)
# Handle stop_ids - return IDs for all special tokens # Handle stop_ids - return IDs for all special tokens
if key == "stop_ids": if key == "stop_ids":
stop_ids = [] stop_ids = []
-3
View File
@@ -1,4 +1,3 @@
from astrai.trainer.optim import Muon
from astrai.trainer.schedule import BaseScheduler, SchedulerFactory from astrai.trainer.schedule import BaseScheduler, SchedulerFactory
from astrai.trainer.strategy import BaseStrategy, StrategyFactory from astrai.trainer.strategy import BaseStrategy, StrategyFactory
from astrai.trainer.train_callback import ( from astrai.trainer.train_callback import (
@@ -10,8 +9,6 @@ from astrai.trainer.trainer import Trainer
__all__ = [ __all__ = [
# Main trainer # Main trainer
"Trainer", "Trainer",
# Optimizer
"Muon",
# Strategy factory # Strategy factory
"StrategyFactory", "StrategyFactory",
"BaseStrategy", "BaseStrategy",
+16 -53
View File
@@ -1,42 +1,25 @@
from typing import Any, Callable, Dict from typing import Dict
import torch import torch
import torch.nn as nn import torch.nn as nn
def _grad_stat( def grad_norm(model: nn.Module, per_param: bool = False) -> float | Dict[str, float]:
model: nn.Module, fn: Callable[[torch.Tensor], Any], default: Any grads = [p.grad.detach() for p in model.parameters() if p.grad is not None]
) -> dict: if not grads:
results = {} return 0.0
total_sq = torch.stack([g.pow(2).sum() for g in grads]).sum()
if per_param:
norms = {}
for name, param in model.named_parameters(): for name, param in model.named_parameters():
results[name] = default
if param.grad is not None: if param.grad is not None:
results[name] = fn(param.grad.data) norms[name] = param.grad.norm(2).item()
return results else:
norms[name] = 0.0
norms["total"] = total_sq.sqrt().item()
def grad_norm(model: nn.Module, norm_type: int = 2) -> Dict[str, float]: return norms
return _grad_stat(model, lambda g: g.norm(norm_type).item(), 0.0) return total_sq.sqrt().item()
def grad_std(model: nn.Module) -> Dict[str, float]:
return _grad_stat(model, lambda g: g.std().item(), 0.0)
def grad_max(model: nn.Module) -> Dict[str, float]:
return _grad_stat(model, lambda g: g.max().item(), -float("inf"))
def grad_min(model: nn.Module) -> Dict[str, float]:
return _grad_stat(model, lambda g: g.min().item(), float("inf"))
def grad_mean(model: nn.Module) -> Dict[str, float]:
return _grad_stat(model, lambda g: g.mean().item(), 0.0)
def grad_nan_num(model: nn.Module) -> Dict[str, int]:
return _grad_stat(model, lambda g: g.isnan().sum().item(), 0)
def ctx_get_loss(ctx): def ctx_get_loss(ctx):
@@ -52,24 +35,4 @@ def ctx_get_val_loss(ctx):
def ctx_get_grad_norm(ctx): def ctx_get_grad_norm(ctx):
return grad_norm(ctx.model) return ctx.grad_norm
def ctx_get_grad_std(ctx):
return grad_std(ctx.model)
def ctx_get_grad_max(ctx):
return grad_max(ctx.model)
def ctx_get_grad_min(ctx):
return grad_min(ctx.model)
def ctx_get_grad_mean(ctx):
return grad_mean(ctx.model)
def ctx_get_grad_nan_num(ctx):
return grad_nan_num(ctx.model)
-143
View File
@@ -1,143 +0,0 @@
import torch
from torch.optim import Optimizer
def _zeropower_via_newtonschulz(G: torch.Tensor, steps: int = 5):
assert G.ndim == 2
X = G
scale = max(1, G.size(0) / G.size(1)) ** 0.5
X = X / (X.norm() + 1e-7) * scale
if steps == 0:
return X
a, b, c = (3.4445, -4.7750, 2.0315)
for _ in range(steps):
A = X @ X.T
B = A @ X
X = a * X + b * B + c * (A @ B)
return X
class Muon(Optimizer):
def __init__(
self,
params,
lr: float = 2e-3,
momentum: float = 0.95,
weight_decay: float = 0.0,
nesterov: bool = True,
ns_steps: int = 5,
adamw_lr: float = None,
adamw_betas: tuple = (0.9, 0.95),
adamw_eps: float = 1e-8,
adamw_wd: float = 0.0,
):
defaults = dict(
lr=lr,
momentum=momentum,
weight_decay=weight_decay,
nesterov=nesterov,
ns_steps=ns_steps,
adamw_lr=adamw_lr if adamw_lr is not None else lr * 0.1,
adamw_betas=adamw_betas,
adamw_eps=adamw_eps,
adamw_wd=adamw_wd,
)
super().__init__(params, defaults)
@torch.no_grad()
def step(self, closure=None):
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()
for group in self.param_groups:
params_2d, params_1d = [], []
grads_2d, grads_1d = [], []
for p in group["params"]:
if p.grad is None:
continue
if p.grad.is_sparse:
raise RuntimeError("Muon does not support sparse gradients")
if p.ndim >= 2:
params_2d.append(p)
grads_2d.append(p.grad)
else:
params_1d.append(p)
grads_1d.append(p.grad)
if params_2d:
self._muon_update_foreach(params_2d, grads_2d, group)
if params_1d:
self._adamw_update_foreach(params_1d, grads_1d, group)
return loss
def _muon_update_foreach(self, params_2d, grads_2d, group):
lr = group["lr"]
momentum = group["momentum"]
wd = group["weight_decay"]
nesterov = group["nesterov"]
ns_steps = group["ns_steps"]
if wd != 0:
torch._foreach_mul_(params_2d, 1 - lr * wd)
if nesterov:
grads_2d = torch._foreach_add(grads_2d, params_2d, alpha=wd)
bufs = []
for p, grad in zip(params_2d, grads_2d):
state = self.state[p]
if "momentum_buffer" not in state:
state["momentum_buffer"] = torch.zeros_like(grad)
bufs.append(state["momentum_buffer"])
torch._foreach_lerp_(bufs, grads_2d, 1 - momentum)
for p, buf in zip(params_2d, bufs):
update = _zeropower_via_newtonschulz(buf, steps=ns_steps)
scale = max(1, p.size(0) / p.size(1)) ** 0.5
p.add_(update, alpha=-lr * scale)
def _adamw_update_foreach(self, params_1d, grads_1d, group):
lr = group["adamw_lr"]
betas = group["adamw_betas"]
eps = group["adamw_eps"]
wd = group["adamw_wd"]
steps: list[int] = []
exp_avgs, exp_avg_sqs = [], []
has_state = []
for p in params_1d:
state = self.state[p]
if not state:
state["step"] = 0
state["exp_avg"] = torch.zeros_like(p)
state["exp_avg_sq"] = torch.zeros_like(p)
has_state.append(False)
else:
has_state.append(True)
state["step"] += 1
steps.append(state["step"])
exp_avgs.append(state["exp_avg"])
exp_avg_sqs.append(state["exp_avg_sq"])
beta1, beta2 = betas
torch._foreach_lerp_(exp_avgs, grads_1d, 1 - beta1)
grads_sq = torch._foreach_mul(grads_1d, grads_1d)
torch._foreach_lerp_(exp_avg_sqs, grads_sq, 1 - beta2)
bias_correction1 = [1 - beta1**s for s in steps]
bias_correction2 = [1 - beta2**s for s in steps]
if wd != 0:
torch._foreach_mul_(params_1d, 1 - lr * wd)
exp_avg_corrected = torch._foreach_div(exp_avgs, bias_correction1)
denom = torch._foreach_div(exp_avg_sqs, bias_correction2)
denom = torch._foreach_sqrt(denom)
torch._foreach_add_(denom, eps)
torch._foreach_addcdiv_(params_1d, exp_avg_corrected, denom, value=-lr)
+13 -7
View File
@@ -53,7 +53,7 @@ class CosineScheduler(BaseScheduler):
optimizer, optimizer,
warmup_steps: int, warmup_steps: int,
lr_decay_steps: int, lr_decay_steps: int,
min_rate: float = 0.05, min_rate: float = 0.01,
last_epoch: int = -1, last_epoch: int = -1,
): ):
self.warmup_steps = warmup_steps self.warmup_steps = warmup_steps
@@ -65,11 +65,15 @@ class CosineScheduler(BaseScheduler):
def get_lr(self) -> List[float]: def get_lr(self) -> List[float]:
# warmup # warmup
if self.last_epoch < self.warmup_steps: if self.last_epoch < self.warmup_steps:
warmup_factor = max(self.min_rate, self.last_epoch / self.warmup_steps) warmup_factor = max(
self.min_rate, self.last_epoch / max(self.warmup_steps, 1)
)
return [base_lr * warmup_factor for base_lr in self.base_lrs] return [base_lr * warmup_factor for base_lr in self.base_lrs]
# cosine decay # cosine decay
decay_progress = (self.last_epoch - self.warmup_steps) / self.lr_decay_steps decay_progress = (self.last_epoch - self.warmup_steps) / max(
self.lr_decay_steps, 1
)
decay_progress = min(decay_progress, 1.0) decay_progress = min(decay_progress, 1.0)
cosine_decay = 0.5 * (1.0 + math.cos(math.pi * decay_progress)) cosine_decay = 0.5 * (1.0 + math.cos(math.pi * decay_progress))
decay_factor = max(self.min_rate, cosine_decay) decay_factor = max(self.min_rate, cosine_decay)
@@ -104,7 +108,7 @@ class SGDRScheduler(BaseScheduler):
optimizer, optimizer,
warmup_steps: int, warmup_steps: int,
cycle_length: int, cycle_length: int,
min_rate: float = 0.05, min_rate: float = 0.01,
t_mult: int = 2, t_mult: int = 2,
last_epoch: int = -1, last_epoch: int = -1,
): ):
@@ -118,7 +122,9 @@ class SGDRScheduler(BaseScheduler):
def get_lr(self): def get_lr(self):
# warmup # warmup
if self.last_epoch < self.warmup_steps: if self.last_epoch < self.warmup_steps:
warmup_factor = max(self.min_rate, self.last_epoch / self.warmup_steps) warmup_factor = max(
self.min_rate, self.last_epoch / max(self.warmup_steps, 1)
)
return [base_lr * warmup_factor for base_lr in self.base_lrs] return [base_lr * warmup_factor for base_lr in self.base_lrs]
# SGDR # SGDR
@@ -182,7 +188,7 @@ class WSDScheduler(BaseScheduler):
warmup_steps: int, warmup_steps: int,
stable_steps: int, stable_steps: int,
decay_steps: int, decay_steps: int,
min_rate: float = 0.0, min_rate: float = 0.01,
last_epoch: int = -1, last_epoch: int = -1,
): ):
self.warmup_steps = warmup_steps self.warmup_steps = warmup_steps
@@ -194,7 +200,7 @@ class WSDScheduler(BaseScheduler):
def get_lr(self) -> List[float]: def get_lr(self) -> List[float]:
if self.last_epoch < self.warmup_steps: if self.last_epoch < self.warmup_steps:
factor = self.last_epoch / max(self.warmup_steps, 1) factor = max(self.min_rate, self.last_epoch / max(self.warmup_steps, 1))
return [base_lr * factor for base_lr in self.base_lrs] return [base_lr * factor for base_lr in self.base_lrs]
offset = self.last_epoch - self.warmup_steps offset = self.last_epoch - self.warmup_steps
+67 -39
View File
@@ -98,7 +98,6 @@ class BaseStrategy(ABC):
self.model = model self.model = model
self.device = device self.device = device
self.executor = kwargs.pop("executor", None) self.executor = kwargs.pop("executor", None)
self.model_fn = kwargs.pop("model_fn", None)
self.extra_kwargs = kwargs self.extra_kwargs = kwargs
@abstractmethod @abstractmethod
@@ -196,7 +195,7 @@ class SFTStrategy(BaseStrategy):
ignore_index = -100 ignore_index = -100
input_mask = make_doc_boundary_mask(position_ids) input_mask = make_doc_boundary_mask(position_ids)
target_ids = target_ids.masked_fill(loss_mask == 0, ignore_index) target_ids = target_ids.masked_fill(~loss_mask, ignore_index)
logits = self.model( logits = self.model(
input_ids=input_ids, position_ids=position_ids, input_mask=input_mask input_ids=input_ids, position_ids=position_ids, input_mask=input_mask
)["logits"] )["logits"]
@@ -223,14 +222,13 @@ class DPOStrategy(BaseStrategy):
self, self,
model: nn.Module, model: nn.Module,
device: str, device: str,
ref_model: nn.Module,
beta: float = 0.1, beta: float = 0.1,
reduction: str = "mean", reduction: str = "sum",
**kwargs, **kwargs,
): ):
super().__init__(model, device, **kwargs) super().__init__(model, device, **kwargs)
self.ref_model = create_ref_model( self.ref_model = ref_model
self.model_fn, self.executor.unwrap_model(model)
).to(device=self.device)
self.beta = beta self.beta = beta
self.reduction = reduction self.reduction = reduction
@@ -267,42 +265,45 @@ class DPOStrategy(BaseStrategy):
class GRPOStrategy(BaseStrategy): class GRPOStrategy(BaseStrategy):
"""Group Relative Policy Optimization strategy. """Group Relative Policy Optimization strategy.
On-policy GRPO following DeepSeek-R1: the policy model is updated while Implements GRPO following DeepSeek-R1 with token-level PPO clipping.
a frozen ref_model stores the old-policy log-probs. ratio = exp(logπ_θ - logπ_ref), Advantages are group-normalized from scalar per-response rewards and
clipped PPO objective. Call ``sync_ref_model()`` after each data-generation round. broadcast across all response tokens. The loss is computed **only on
response tokens** prompt tokens are masked out.
Three model roles are distinguished:
* **Policy** ``self.model`` the model being trained.
* **Old policy** ``self.old_model`` the behaviour policy that generated
the responses. Used for the importance sampling ratio
``ρ = π_θ / π_old``. Synced externally after each data-generation round.
* **Reference model** ``self.ref_model`` a frozen copy of the initial
policy (typically the SFT checkpoint) used **only** for the KL
regularisation term. It is never updated during training.
""" """
def __init__( def __init__(
self, self,
model: nn.Module, model: nn.Module,
device: str, device: str,
old_model: nn.Module,
ref_model: nn.Module,
clip_eps: float = 0.2, clip_eps: float = 0.2,
kl_coef: float = 0.01, kl_coef: float = 0.01,
group_size: int = 4, group_size: int = 4,
reduction: str = "mean",
sync_interval: int = 200,
**kwargs, **kwargs,
): ):
super().__init__(model, device, **kwargs) super().__init__(model, device, **kwargs)
self.ref_model = create_ref_model( self.old_model = old_model
self.model_fn, self.executor.unwrap_model(model) self.ref_model = ref_model
).to(device=self.device)
self.clip_eps = clip_eps self.clip_eps = clip_eps
self.kl_coef = kl_coef self.kl_coef = kl_coef
self.group_size = group_size self.group_size = group_size
self.reduction = reduction
self.sync_interval = sync_interval
self._step = 0
def sync_ref_model(self): def sync_old_model(self):
"""Copy current model weights to ref model.""" """Copy current policy weights to old model."""
self.ref_model.load_state_dict(self.executor.unwrap_model(self.model)) self.old_model.load_state_dict(self.executor.unwrap_model(self.model))
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor: def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
self._step += 1
if self._step % self.sync_interval == 0:
self.sync_ref_model()
batch = move_to_device(batch, self.device) batch = move_to_device(batch, self.device)
prompts = batch["prompts"] prompts = batch["prompts"]
responses = batch["responses"] responses = batch["responses"]
@@ -313,33 +314,60 @@ class GRPOStrategy(BaseStrategy):
responses_flat = responses.view(-1, response_len) responses_flat = responses.view(-1, response_len)
masks_flat = masks.view(-1, response_len) masks_flat = masks.view(-1, response_len)
prompt_expanded = prompts.unsqueeze(1).repeat(1, group_size, 1).flatten(0, 1) prompt_expanded = prompts.unsqueeze(1).repeat(1, group_size, 1).flatten(0, 1)
prompt_len = prompt_expanded.size(1)
full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-1) full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-1)
full_masks = torch.cat([torch.ones_like(prompt_expanded), masks_flat], dim=-1) # Prompt tokens are masked out (0) so logprobs are computed only for
# response tokens. get_logprobs shifts the mask by one position, so
log_probs_policy = get_logprobs( # the first response token's logprob (predicted from the last prompt
self.model, full_sequences, full_masks, self.reduction # token) is correctly included.
) full_masks = torch.cat([torch.zeros_like(prompt_expanded), masks_flat], dim=-1)
log_probs_policy = log_probs_policy.view(batch_size, group_size)
# get_logprobs returns [B*G, S-1] (S = prompt_len + response_len).
# Response token logprobs occupy the last ``response_len`` positions
# (the first response token is predicted from the last prompt token).
token_log_probs_policy = get_logprobs(
self.model, full_sequences, full_masks, "none"
)[:, prompt_len - 1 :]
with torch.no_grad(): with torch.no_grad():
log_probs_ref = get_logprobs( token_log_probs_old = get_logprobs(
self.ref_model, full_sequences, full_masks, self.reduction self.old_model, full_sequences, full_masks, "none"
) )[:, prompt_len - 1 :]
log_probs_ref = log_probs_ref.view(batch_size, group_size) token_log_probs_ref = get_logprobs(
self.ref_model, full_sequences, full_masks, "none"
)[:, prompt_len - 1 :]
eps = torch.finfo(log_probs_policy.dtype).eps # Reshape to [B, G, response_len]
token_log_probs_policy = token_log_probs_policy.view(batch_size, group_size, -1)
token_log_probs_old = token_log_probs_old.view(batch_size, group_size, -1)
token_log_probs_ref = token_log_probs_ref.view(batch_size, group_size, -1)
token_masks = masks_flat.view(batch_size, group_size, -1).float()
# Group-normalized advantages from scalar per-response rewards.
eps = 1e-8
mean = rewards.mean(dim=-1, keepdim=True) mean = rewards.mean(dim=-1, keepdim=True)
std = rewards.std(dim=-1, keepdim=True) std = rewards.std(dim=-1, keepdim=True, unbiased=False)
advantages = (rewards - mean) / (std + eps) advantages = (rewards - mean) / (std + eps)
# Broadcast scalar advantage to every response token: [B, G, 1]
advantages = advantages.unsqueeze(-1)
ratio = torch.exp(log_probs_policy - log_probs_ref) # Token-level ratio (π_θ / π_old) and PPO clipping.
log_ratio = token_log_probs_policy - token_log_probs_old
ratio = torch.exp(log_ratio)
surr1 = ratio * advantages surr1 = ratio * advantages
surr2 = torch.clamp(ratio, 1 - self.clip_eps, 1 + self.clip_eps) * advantages surr2 = torch.clamp(ratio, 1 - self.clip_eps, 1 + self.clip_eps) * advantages
per_token_policy_loss = -torch.min(surr1, surr2)
token_count = token_masks.sum().clamp(min=1.0)
policy_loss = (per_token_policy_loss * token_masks).sum() / token_count
# KL penalty to frozen reference model with k1 estimator (non-negative):
# k1 = π_ref / π_θ - log(π_ref / π_θ) - 1, where π_ref / π_θ = exp(log_ref - log_policy).
log_ref_ratio = token_log_probs_ref - token_log_probs_policy
r = torch.exp(log_ref_ratio)
kl_per_token = r - torch.log(r + eps) - 1.0
kl_penalty = self.kl_coef * (kl_per_token * token_masks).sum() / token_count
policy_loss = -torch.min(surr1, surr2).mean()
kl_penalty = self.kl_coef * (log_probs_policy - log_probs_ref).square().mean()
total_loss = policy_loss + kl_penalty total_loss = policy_loss + kl_penalty
return total_loss return total_loss
+93 -89
View File
@@ -9,21 +9,15 @@ from typing import IO, Callable, List, Optional, Protocol, runtime_checkable
import torch import torch
import torch.distributed as dist import torch.distributed as dist
import torch.nn as nn import torch.nn as nn
from torch.nn.utils import clip_grad_norm_
from torch.utils.checkpoint import checkpoint as torch_checkpoint from torch.utils.checkpoint import checkpoint as torch_checkpoint
from tqdm import tqdm from tqdm import tqdm
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.parallel import only_on_rank from astrai.parallel import only_on_rank
from astrai.parallel.setup import get_current_device, get_rank from astrai.parallel.setup import get_current_device
from astrai.serialization import Checkpoint from astrai.serialization import Checkpoint
from astrai.trainer.metric_util import ( from astrai.trainer.metric_util import (
ctx_get_grad_max,
ctx_get_grad_mean,
ctx_get_grad_min,
ctx_get_grad_nan_num,
ctx_get_grad_norm, ctx_get_grad_norm,
ctx_get_grad_std,
ctx_get_loss, ctx_get_loss,
ctx_get_lr, ctx_get_lr,
ctx_get_val_loss, ctx_get_val_loss,
@@ -86,7 +80,9 @@ class GradientClippingCallback(TrainCallback):
self.max_grad_norm = max_grad_norm self.max_grad_norm = max_grad_norm
def on_optimizer_step(self, context: TrainContext): def on_optimizer_step(self, context: TrainContext):
clip_grad_norm_(context.model.parameters(), self.max_grad_norm) context.grad_norm = context.executor.clip_grad_norm(
context.model, self.max_grad_norm
)
@CallbackFactory.register("gradient_checkpointing") @CallbackFactory.register("gradient_checkpointing")
@@ -143,34 +139,38 @@ class CheckpointCallback(TrainCallback):
self.interval = interval self.interval = interval
self.weight_only = weight_only self.weight_only = weight_only
self.save_extra_fn = save_extra_fn or CheckpointCallback.save_extra self.save_extra_fn = save_extra_fn or CheckpointCallback.save_extra
self.last_ckpt_iter = 0 self.last_ckpt_step = None
def on_train_begin(self, context: TrainContext):
self.last_ckpt_step = context.optimizer_step
def _save_checkpoint(self, context: TrainContext): def _save_checkpoint(self, context: TrainContext):
state_dict = context.executor.unwrap_model(context.model) self.last_ckpt_step = context.optimizer_step
self.last_ckpt_iter = context.iteration
if get_rank() == 0: with context.executor.checkpoint_context(context.model) as state_dict:
if state_dict is not None:
save_path = os.path.join( save_path = os.path.join(
self.save_dir, f"epoch_{context.epoch}_iter_{context.iteration}" self.save_dir,
f"epoch_{context.epoch}_step_{context.optimizer_step}",
) )
extra = self.save_extra_fn(context) extra = self.save_extra_fn(context)
meta = context.config.to_dict() meta = context.config.to_dict()
context.checkpoint = Checkpoint( context.checkpoint = Checkpoint(
state_dict=state_dict, state_dict=state_dict,
epoch=context.epoch, epoch=context.epoch,
iteration=context.iteration, consumed_samples=context.consumed_samples,
config=context.model_config,
extra=extra, extra=extra,
meta=meta, meta=meta,
config=context.model_config,
) )
context.checkpoint.save(save_path) context.checkpoint.save(save_path)
def on_batch_end(self, context: TrainContext): def on_batch_end(self, context: TrainContext):
if context.iteration - self.last_ckpt_iter >= self.interval: if context.optimizer_step - self.last_ckpt_step >= self.interval:
self._save_checkpoint(context) self._save_checkpoint(context)
def on_train_end(self, context: TrainContext): def on_train_end(self, context: TrainContext):
if context.iteration != self.last_ckpt_iter: if context.optimizer_step != self.last_ckpt_step:
self._save_checkpoint(context) self._save_checkpoint(context)
def on_error(self, context: TrainContext): def on_error(self, context: TrainContext):
@@ -202,19 +202,23 @@ class ProgressBarCallback(TrainCallback):
@only_on_rank(0) @only_on_rank(0)
def on_epoch_begin(self, context: TrainContext): def on_epoch_begin(self, context: TrainContext):
total_steps = len(context.dataloader) // context.executor.grad_accum_steps
self.progress_bar = tqdm( self.progress_bar = tqdm(
context.dataloader, total=total_steps,
desc=f"Epoch {context.epoch + 1}/{self.num_epoch}", desc=f"Epoch {context.epoch + 1}/{self.num_epoch}",
dynamic_ncols=True, dynamic_ncols=True,
file=self.file or sys.stdout, file=self.file or sys.stdout,
) )
@only_on_rank(0) @only_on_rank(0)
def on_batch_end(self, context: TrainContext): def on_optimizer_step(self, context: TrainContext):
postfix = { postfix = {
"step": f"{context.optimizer_step:d}",
"loss": f"{context.loss:.4f}", "loss": f"{context.loss:.4f}",
"lr": f"{context.optimizer.param_groups[-1]['lr']:.2e}", "lr": f"{context.optimizer.param_groups[-1]['lr']:.2e}",
} }
if context.grad_norm is not None:
postfix["grad_norm"] = f"{context.grad_norm:.2f}"
if context.val_loss is not None: if context.val_loss is not None:
postfix["val_loss"] = f"{context.val_loss:.4f}" postfix["val_loss"] = f"{context.val_loss:.4f}"
self.progress_bar.set_postfix(postfix) self.progress_bar.set_postfix(postfix)
@@ -227,19 +231,20 @@ class ProgressBarCallback(TrainCallback):
self.progress_bar.close() self.progress_bar.close()
@CallbackFactory.register("metric_logger") @CallbackFactory.register("metric")
class MetricLoggerCallback(TrainCallback): class MetricCallback(TrainCallback):
def __init__( def __init__(
self, self,
log_dir: str, log_dir: str,
save_interval: int, save_interval: int,
log_interval: int = 10,
metrics: List[str] = None, metrics: List[str] = None,
val_step: int = 0,
): ):
self.last_log_iter = 0 self.last_log_flush_step = None
self.save_interval = save_interval self.save_interval = save_interval
self.log_interval = log_interval
self.metrics = metrics or ["loss", "lr"] self.metrics = metrics or ["loss", "lr"]
self.val_step = val_step
self._next_val_step = 0
self.log_dir = Path(log_dir) if log_dir else Path.cwd() / "logs" self.log_dir = Path(log_dir) if log_dir else Path.cwd() / "logs"
self.log_dir.mkdir(parents=True, exist_ok=True) self.log_dir.mkdir(parents=True, exist_ok=True)
@@ -251,58 +256,28 @@ class MetricLoggerCallback(TrainCallback):
"lr": ctx_get_lr, "lr": ctx_get_lr,
"val_loss": ctx_get_val_loss, "val_loss": ctx_get_val_loss,
"grad_norm": ctx_get_grad_norm, "grad_norm": ctx_get_grad_norm,
"grad_std": ctx_get_grad_std,
"grad_max": ctx_get_grad_max,
"grad_min": ctx_get_grad_min,
"grad_mean": ctx_get_grad_mean,
"grad_nan_num": ctx_get_grad_nan_num,
} }
def _get_log_data(self, context: TrainContext): def _metrics(self, context: TrainContext, names):
data = { return {
m: self._metric_funcs[m](context)
for m in names
if self._metric_funcs[m](context) is not None
}
@only_on_rank(0)
def _append(self, event_type: str, context: TrainContext, **extra):
entry = {
"type": event_type,
"timestamp": time.strftime("%Y-%m-%dT%H:%M:%S"), "timestamp": time.strftime("%Y-%m-%dT%H:%M:%S"),
"epoch": context.epoch, "epoch": context.epoch,
"iter": context.iteration, "step": context.optimizer_step,
"consumed_samples": context.consumed_samples,
**extra,
} }
for m in self.metrics: self.log_cache.append(entry)
val = self._metric_funcs[m](context)
if val is not None:
data[m] = val
return data
@only_on_rank(0) def _run_validation(self, context: TrainContext) -> float:
def _add_log(self, log_data):
self.log_cache.append(log_data)
@only_on_rank(0)
def _save_log(self, epoch, iter):
log_file = self.log_dir / f"epoch_{epoch}_iter_{iter}_metric.jsonl"
log_file.parent.mkdir(parents=True, exist_ok=True)
with open(log_file, "w") as f:
for log in self.log_cache:
f.write(json.dumps(log) + "\n")
def on_batch_end(self, context):
if context.iteration % self.log_interval == 0:
log_data = self._get_log_data(context)
self._add_log(log_data)
if context.iteration - self.last_log_iter >= self.save_interval:
self._save_log(context.epoch, context.iteration)
self.last_log_iter = context.iteration
def on_train_end(self, context):
if context.iteration != self.last_log_iter:
self._save_log(context.epoch, context.iteration)
def on_error(self, context):
self._save_log(context.epoch, context.iteration)
@CallbackFactory.register("validation")
class ValidationCallback(TrainCallback):
def _run_validation(self, context: TrainContext):
context.model.eval() context.model.eval()
total_loss = 0.0 total_loss = 0.0
@@ -314,27 +289,56 @@ class ValidationCallback(TrainCallback):
total_loss += loss.item() total_loss += loss.item()
num_batches += 1 num_batches += 1
if context.world_size > 1 and dist.is_initialized():
stats = torch.tensor(
[total_loss, float(num_batches)], device=get_current_device()
)
dist.all_reduce(stats, op=dist.ReduceOp.SUM)
avg_loss = (stats[0] / stats[1]).item()
else:
avg_loss = total_loss / max(num_batches, 1) avg_loss = total_loss / max(num_batches, 1)
if context.world_size > 1 and dist.is_initialized():
loss_tensor = torch.tensor([avg_loss], device=get_current_device())
dist.all_reduce(loss_tensor, op=dist.ReduceOp.AVG)
avg_loss = loss_tensor.item()
context.val_loss = avg_loss
context.model.train() context.model.train()
return avg_loss
step_count = context.iteration // context.config.grad_accum_steps def on_train_begin(self, context: TrainContext):
logger.info( self.last_log_flush_step = context.optimizer_step
f"Epoch {context.epoch + 1}, Step {step_count}, Val Loss: {avg_loss:.4f}"
)
def on_optimizer_step(self, context: TrainContext): @only_on_rank(0)
if context.val_dataloader is None: def _flush(self, epoch, step):
return log_file = self.log_dir / f"epoch_{epoch}_step_{step}_metric.jsonl"
cfg = context.config log_file.parent.mkdir(parents=True, exist_ok=True)
if cfg.val_step <= 0: with open(log_file, "w") as f:
return for log in self.log_cache:
step_count = context.iteration // cfg.grad_accum_steps f.write(json.dumps(log) + "\n")
if step_count % cfg.val_step == 0:
self._run_validation(context) def on_optimizer_step(self, context):
if (
context.val_dataloader is not None
and self.val_step > 0
and context.optimizer_step >= self._next_val_step
):
context.val_loss = self._run_validation(context)
self._next_val_step = context.optimizer_step + self.val_step
self._append("validation", context, val_loss=context.val_loss)
step_metrics = [m for m in self.metrics if m != "val_loss"]
self._append("step", context, **self._metrics(context, step_metrics))
if context.optimizer_step - self.last_log_flush_step >= self.save_interval:
self._flush(context.epoch, context.optimizer_step)
self.last_log_flush_step = context.optimizer_step
def on_epoch_end(self, context):
self._append("epoch", context)
def on_train_end(self, context):
if (
self.last_log_flush_step is None
or context.optimizer_step != self.last_log_flush_step
):
self._flush(context.epoch, context.optimizer_step)
self.last_log_flush_step = context.optimizer_step
def on_error(self, context):
self._flush(context.epoch, context.optimizer_step)
+56 -17
View File
@@ -7,13 +7,13 @@ import torch.nn as nn
from torch.utils.data import DataLoader, random_split from torch.utils.data import DataLoader, random_split
from astrai.config.train_config import TrainConfig from astrai.config.train_config import TrainConfig
from astrai.dataset import ResumableDistributedSampler from astrai.dataset import RDSampler
from astrai.model.components.lora import inject_lora from astrai.model.components.lora import inject_lora
from astrai.parallel.executor import BaseExecutor, ExecutorFactory from astrai.parallel.executor import BaseExecutor, ExecutorFactory
from astrai.parallel.setup import get_current_device, get_rank, get_world_size from astrai.parallel.setup import get_current_device, get_rank, get_world_size
from astrai.protocols import OptimizerProtocol, SchedulerProtocol from astrai.protocols import OptimizerProtocol, SchedulerProtocol
from astrai.serialization import Checkpoint, load_json from astrai.serialization import Checkpoint, load_json
from astrai.trainer.strategy import BaseStrategy, StrategyFactory from astrai.trainer.strategy import BaseStrategy, StrategyFactory, create_ref_model
@dataclass @dataclass
@@ -29,8 +29,9 @@ class TrainContext:
executor: BaseExecutor = field(default=None) executor: BaseExecutor = field(default=None)
epoch: int = field(default=0) epoch: int = field(default=0)
iteration: int = field(default=0) consumed_samples: int = field(default=0)
loss: float = field(default=0.0) loss: float = field(default=0.0)
grad_norm: Optional[float] = field(default=None)
val_dataloader: Optional[DataLoader] = field(default=None) val_dataloader: Optional[DataLoader] = field(default=None)
val_loss: Optional[float] = field(default=None) val_loss: Optional[float] = field(default=None)
@@ -38,6 +39,14 @@ class TrainContext:
rank: int = field(default=0) rank: int = field(default=0)
kwargs: Dict[str, Any] = field(default_factory=dict) kwargs: Dict[str, Any] = field(default_factory=dict)
@property
def optimizer_step(self) -> int:
return self.consumed_samples // (
self.config.batch_per_device
* self.world_size
* self.config.grad_accum_steps
)
class TrainContextBuilder: class TrainContextBuilder:
def __init__( def __init__(
@@ -45,10 +54,12 @@ class TrainContextBuilder:
config: TrainConfig, config: TrainConfig,
): ):
self.config = config self.config = config
self._resume_dir: Optional[str] = None self._param_path: Optional[str] = None
self._resume: bool = False
def with_resume_dir(self, resume_dir: Optional[str]) -> Self: def with_param_path(self, param_path: Optional[str], resume: bool = False) -> Self:
self._resume_dir = resume_dir self._param_path = param_path
self._resume = resume
return self return self
def build(self) -> TrainContext: def build(self) -> TrainContext:
@@ -63,11 +74,10 @@ class TrainContextBuilder:
model = cfg.model_fn() model = cfg.model_fn()
model = model.to(device=device) model = model.to(device=device)
model.embed_tokens.neftune_noise_alpha = cfg.neftune_alpha
model_config = {} model_config = {}
if self._resume_dir: if self._param_path:
config_path = Path(self._resume_dir) / "config.json" config_path = Path(self._param_path) / "config.json"
if config_path.exists(): if config_path.exists():
model_config = load_json(config_path) model_config = load_json(config_path)
@@ -83,14 +93,28 @@ class TrainContextBuilder:
executor=executor, executor=executor,
) )
if self._resume_dir: if self._param_path:
checkpoint = Checkpoint.load_any(self._resume_dir) checkpoint = Checkpoint.load_any(self._param_path)
if checkpoint is not None: if checkpoint is not None:
model.load_state_dict(checkpoint.state_dict, strict=False) model.load_state_dict(checkpoint.state_dict, strict=False)
if checkpoint.config: if checkpoint.config:
context.model_config = checkpoint.config context.model_config = checkpoint.config
if self._resume:
context.epoch = checkpoint.epoch or cfg.start_epoch context.epoch = checkpoint.epoch or cfg.start_epoch
context.iteration = checkpoint.iteration or cfg.start_batch 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 context.checkpoint = checkpoint
if cfg.lora is not None: if cfg.lora is not None:
@@ -116,8 +140,8 @@ class TrainContextBuilder:
cfg.dataset, [n_train, n_val], generator=generator cfg.dataset, [n_train, n_val], generator=generator
) )
sampler_offset = context.iteration * cfg.batch_per_device sampler_offset = context.consumed_samples // context.world_size
sampler = ResumableDistributedSampler( sampler = RDSampler(
data_source=train_dataset, data_source=train_dataset,
start_epoch=context.epoch, start_epoch=context.epoch,
start_iter=sampler_offset, start_iter=sampler_offset,
@@ -130,10 +154,11 @@ class TrainContextBuilder:
num_workers=cfg.num_workers, num_workers=cfg.num_workers,
pin_memory=cfg.pin_memory, pin_memory=cfg.pin_memory,
prefetch_factor=cfg.prefetch_factor, prefetch_factor=cfg.prefetch_factor,
collate_fn=cfg.collate_fn,
) )
if val_dataset is not None: if val_dataset is not None:
val_sampler = ResumableDistributedSampler( val_sampler = RDSampler(
data_source=val_dataset, data_source=val_dataset,
start_epoch=0, start_epoch=0,
start_iter=0, start_iter=0,
@@ -147,6 +172,7 @@ class TrainContextBuilder:
num_workers=cfg.num_workers, num_workers=cfg.num_workers,
pin_memory=cfg.pin_memory, pin_memory=cfg.pin_memory,
prefetch_factor=cfg.prefetch_factor, prefetch_factor=cfg.prefetch_factor,
collate_fn=cfg.collate_fn,
) )
context.model, context.optimizer, context.dataloader, context.scheduler = ( context.model, context.optimizer, context.dataloader, context.scheduler = (
@@ -166,13 +192,26 @@ class TrainContextBuilder:
if obj is not None: if obj is not None:
obj.load_state_dict(extra[name]) obj.load_state_dict(extra[name])
strategy_kwargs = dict(cfg.extra_kwargs)
if cfg.strategy in ("dpo", "grpo"):
ref_model = create_ref_model(
cfg.model_fn, executor.unwrap_model(context.model)
).to(device=device)
strategy_kwargs["ref_model"] = ref_model
if cfg.strategy == "grpo":
old_model = create_ref_model(
cfg.model_fn, executor.unwrap_model(context.model)
).to(device=device)
strategy_kwargs["old_model"] = old_model
context.strategy = StrategyFactory.create( context.strategy = StrategyFactory.create(
cfg.strategy, cfg.strategy,
model=context.model, model=context.model,
device=device, device=device,
executor=executor, executor=executor,
model_fn=cfg.model_fn, **strategy_kwargs,
**cfg.extra_kwargs,
) )
return context return context
+12 -8
View File
@@ -35,15 +35,14 @@ class Trainer:
cfg.ckpt_interval, cfg.ckpt_interval,
), ),
CallbackFactory.create( CallbackFactory.create(
"metric_logger", "metric",
log_dir=cfg.log_dir, log_dir=cfg.log_dir,
save_interval=cfg.ckpt_interval, save_interval=cfg.ckpt_interval,
log_interval=cfg.log_interval,
metrics=cfg.metrics, metrics=cfg.metrics,
val_step=cfg.val_step,
), ),
CallbackFactory.create("progress_bar", cfg.n_epoch), CallbackFactory.create("progress_bar", cfg.n_epoch),
CallbackFactory.create("gradient_clipping", cfg.max_grad_norm), CallbackFactory.create("gradient_clipping", cfg.max_grad_norm),
CallbackFactory.create("validation"),
] ]
return callbacks return callbacks
@@ -53,9 +52,11 @@ class Trainer:
if method: if method:
method(context) method(context)
def _trainer_loop(self, resume_dir: Optional[str] = None): def _trainer_loop(self, param_path: Optional[str] = None, resume: bool = False):
context = ( context = (
TrainContextBuilder(self.train_config).with_resume_dir(resume_dir).build() TrainContextBuilder(self.train_config)
.with_param_path(param_path, resume=resume)
.build()
) )
executor = context.executor executor = context.executor
self._call_callbacks("on_train_begin", context) self._call_callbacks("on_train_begin", context)
@@ -74,7 +75,9 @@ class Trainer:
context.loss = loss.item() context.loss = loss.item()
stand_loss = loss / executor.grad_accum_steps stand_loss = loss / executor.grad_accum_steps
executor.backward(stand_loss) executor.backward(stand_loss)
context.iteration += 1 context.consumed_samples += (
context.config.batch_per_device * context.world_size
)
self._call_callbacks("on_batch_end", context) self._call_callbacks("on_batch_end", context)
if executor.sync_gradients: if executor.sync_gradients:
@@ -94,7 +97,7 @@ class Trainer:
finally: finally:
self._call_callbacks("on_train_end", context) self._call_callbacks("on_train_end", context)
def train(self, resume_dir: Optional[str] = None): def train(self, param_path: Optional[str] = None, resume: bool = False):
cfg = self.train_config cfg = self.train_config
spawn_parallel_fn( spawn_parallel_fn(
self._trainer_loop, self._trainer_loop,
@@ -104,5 +107,6 @@ class Trainer:
master_port=cfg.master_port, master_port=cfg.master_port,
device_type=cfg.device_type, device_type=cfg.device_type,
start_method=cfg.start_method, start_method=cfg.start_method,
resume_dir=resume_dir, param_path=param_path,
resume=resume,
) )
+2
View File
@@ -0,0 +1,2 @@
# Source directory for CUDA kernels — build-time only.
# Compiled .so files live in astrAI/_ext/.
+48
View File
@@ -0,0 +1,48 @@
from pathlib import Path
def _arch_flags() -> list[str]:
import torch
if torch.cuda.is_available():
cap = torch.cuda.get_device_capability()
else:
cap = (8, 0)
ver = f"{cap[0]}{cap[1]}"
flags = [f"-gencode=arch=compute_{ver},code=sm_{ver}"]
# tensor-core mma path (mma.sync.m16n8k16.bf16) requires sm_80+; decide the
# kernel dispatch at build time via this define rather than at runtime.
if cap[0] < 8:
flags.append("-DASTRAI_NO_MMA")
return flags
_kernels_dir = Path("csrc/kernels")
REGISTRY: dict[str, dict] = {}
CXX_FLAGS = ["-O3", "-funroll-loops"]
NVCC_FLAGS = [
"-O3",
"--expt-relaxed-constexpr",
"--use_fast_math",
"--ptxas-options=-O3,-v",
"--extra-device-vectorization",
"--threads=8",
]
def register(name: str, sources: list[str] | None = None, **kwargs):
if sources is None:
sources = [str(_kernels_dir / f"{name}.cu")]
REGISTRY[name] = {
"sources": sources,
"cxx_flags": [*CXX_FLAGS],
"nvcc_flags": [*NVCC_FLAGS, *_arch_flags()],
"extra_link_args": kwargs.pop("extra_link_args", []),
**kwargs,
}
register("attn_decode")
register("attn_prefill")
register("attn_paged_decode")
+68
View File
@@ -0,0 +1,68 @@
#pragma once
template<typename T, typename AT = float>
struct AttentionParams {
int batch;
int q_head;
int kv_head;
int q_len;
int kv_len;
int head_dim;
int use_mask;
int causal_offset; // -1 = non-causal; >=0 = absolute position of first Q token
int num_splits;
float scale;
// Q strides (element offsets for each dim — layout-agnostic)
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
// KV strides (K and V share the same layout — only base pointers differ)
int kv_stride_b, kv_stride_h, kv_stride_l, kv_stride_d;
// Mask: 2D [batch, kv_len] (mask_q_stride=0) or 3D [batch, q_len, kv_len]
int mask_b_stride; // = kv_len (both 2D and 3D)
int mask_q_stride; // 2D: 0 (all q rows share); 3D: kv_len
const T* __restrict__ q;
const T* __restrict__ k;
const T* __restrict__ v;
const bool* __restrict__ mask;
T* __restrict__ o;
AT* __restrict__ o_part;
AT* __restrict__ ml_part;
};
template<typename T, typename AT = float>
struct PagedAttentionParams {
int batch;
int q_head;
int kv_head;
int q_len;
int kv_len;
int head_dim;
int use_mask;
int causal_offset;
float scale;
int num_splits;
int page_size;
int max_pages;
// Q strides (layout-agnostic)
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
// Mask strides (2D or 3D)
int mask_b_stride;
int mask_q_stride;
const T* __restrict__ q;
const T* __restrict__ k_cache;
const T* __restrict__ v_cache;
const bool* __restrict__ mask;
const int64_t* __restrict__ page_table;
T* __restrict__ o;
AT* __restrict__ o_part;
AT* __restrict__ ml_part;
};
+82
View File
@@ -0,0 +1,82 @@
#include "attn_decode_split_kv.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,
torch::Tensor v,
c10::optional<torch::Tensor> mask,
int64_t causal_offset,
double scale,
int64_t layout
) {
AttentionParams<bf16> p;
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
TORCH_CHECK(p.q_len == 1, "Q seq_len must be 1");
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
// 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();
DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p);
return O;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("attn_decode", &attn_decode,
py::arg("q"),
py::arg("k"),
py::arg("v"),
py::arg("mask") = py::none(),
py::arg("causal_offset") = -1,
py::arg("scale") = 0.0,
py::arg("layout") = 0,
"GQA decode (tensor-core head-packing on sm_80+, scalar fallback)");
}
+132
View File
@@ -0,0 +1,132 @@
#pragma once
#include <cuda_bf16.h>
#include <float.h>
#include "attn_common.h"
using bf16 = __nv_bfloat16;
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;
}
__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;
int split = blockIdx.z;
int group_size = blockDim.y;
int q_head = kv_head * group_size + threadIdx.y;
int lane = threadIdx.x;
int hd_per_thread = p.head_dim / 32;
// Q: [batch, q_head, q_len=1, head_dim] — stride-based
float q_reg[8];
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
+ lane * hd_per_thread * p.q_stride_d;
for (int i = 0; i < hd_per_thread; i++)
q_reg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
// KV: [batch, kv_head, kv_len, head_dim] — stride-based base
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
int mask_base = batch * p.mask_b_stride;
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
extern __shared__ __align__(16) bf16 k_smem[];
// Split-KV: each split processes a contiguous subset of chunks
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
int chunks_per_split = (chunks_total + p.num_splits - 1) / p.num_splits;
int ch_begin = split * chunks_per_split;
int ch_end = min(chunks_total, ch_begin + chunks_per_split);
for (int ci = ch_begin; ci < ch_end; ci++) {
int chunk_start = ci * DC_CHUNK;
int this_chunk = min(DC_CHUNK, p.kv_len - chunk_start);
// Load K into shared memory (gather from strided global)
int total = this_chunk * p.head_dim;
for (int i = threadIdx.y * 32 + lane; i < total; i += blockDim.x * blockDim.y) {
int s = i / p.head_dim;
int d_dim = i % p.head_dim;
int kv_idx = chunk_start + s;
int g_off = kv_base + kv_idx * p.kv_stride_l + d_dim * p.kv_stride_d;
k_smem[i] = p.k[g_off];
}
__syncthreads();
for (int s = 0; s < this_chunk; s++) {
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 = 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;
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;
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;
m = new_m;
}
__syncthreads();
}
// ---- 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;
int d0 = lane * hd_per_thread;
for (int i = 0; i < hd_per_thread; i++) {
int dd = d0 + i;
p.o_part[slot * p.head_dim + dd] = acc_reg[i];
}
if (lane == 0) {
p.ml_part[slot * 2] = m;
p.ml_part[slot * 2 + 1] = d;
}
}
// 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;
if (d >= p.head_dim) return;
int batch = bh / p.q_head;
int q_head = bh % p.q_head;
size_t split_base = (size_t)bh * p.num_splits;
const float* mlp = p.ml_part + split_base * 2;
const float* op = p.o_part + split_base * p.head_dim;
float m = -FLT_MAX, l = 0.0f, acc = 0.0f;
for (int s = 0; s < p.num_splits; s++) {
float mi = mlp[s * 2];
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;
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);
}
+176
View File
@@ -0,0 +1,176 @@
#pragma once
#include <cfloat>
#include <cuda_bf16.h>
#include "attn_common.h"
#include "attn_mma_utils.cuh"
using bf16 = __nv_bfloat16;
// Split-K (FlashDecoding) tensor-core decode via GQA head-packing.
//
// Decode has q_len == 1, so S = q @ K^T is a GEMV per head — no tensor-core
// work on its own. But GQA gives us G = q_head / kv_head query heads that all
// share one kv_head. We pack those G heads into the M=16 rows of
// mma.sync.m16n8k16, turning G independent GEMVs into a single GEMM that
// reuses each loaded K/V tile across all G heads (K/V load is the decode
// bottleneck, so the reuse is the win, not the flops). The KV sequence is
// partitioned across gridDim.z blocks so that a decode with only
// batch*kv_head independent tasks can fill all SMs. Each (batch, kv_head,
// split) block computes an UN-normalised partial (Oacc, m, l) over its KV
// slice; the combine kernel below reduces across splits. Fixes the "grid too
// small" bottleneck (0.04 waves/SM → many blocks) for long-context,
// small-batch decode.
template <int HEAD_DIM, int BC, int STAGES = 2>
__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 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];
// ---- 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);
float Oacc[DN8][4];
#pragma unroll
for (int j = 0; j < DN8; j++)
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
// KV: stride-based base — [batch, kv_head, kv_len, head_dim]
const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
const int tiles_total = (p.kv_len + BC - 1) / BC;
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
const int ti_begin = split * tiles_per_split;
const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
const int has_mask = p.use_mask && p.mask;
// ---- Load tile lambda: predicated cp.async, unified full/partial ----
auto load_tile = [&](int ti, int buf) {
int kv0 = ti * BC;
bf16* dK = sK + buf * BC * LD;
bf16* dV = sV + buf * BC * LD;
#pragma unroll
for (int i = lane * VEC; i < TOTAL; i += 32 * VEC) {
int r = i / HEAD_DIM, d = i % HEAD_DIM;
int kc = kv0 + r;
bool valid = kc < p.kv_len;
int off = r * LD + swiz_col(d, r, SWIZ_MASK);
// KV stride-based: contiguous within head_dim (stride_d == 1 typically)
int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d;
cp_async_16_pred(&dK[off], &p.k[g_off], valid);
cp_async_16_pred(&dV[off], &p.v[g_off], valid);
}
cp_async_commit();
};
// ---- Prologue: issue first tile load ----
if (ti_begin < ti_end) {
load_tile(ti_begin, 0);
}
for (int ti = ti_begin; ti < ti_end; ti++) {
constexpr int BUF_MASK = (STAGES > 1) ? (STAGES - 1) : 0;
int buf = (ti - ti_begin) & BUF_MASK;
// Wait for current tile, then issue next tile's prefetch (overlaps
// with this tile's compute). Single syncwarp covers both hazards.
// When STAGES==1, no prefetch — load happens at end of prior iter.
cp_async_wait_group<0>();
__syncwarp();
if constexpr (STAGES > 1) {
if (ti + 1 < ti_end)
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
}
const bf16* bK = sK + buf * BC * LD;
const bf16* bV = sV + buf * BC * LD;
int kv0 = ti * BC;
float Sacc[NC8][4];
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
#pragma unroll
for (int n8 = 0; n8 < NC8; n8++)
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
// 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);
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc);
__syncwarp();
if constexpr (STAGES == 1) {
if (ti + 1 < ti_end)
load_tile(ti + 1, 0);
}
}
// ---- write UN-normalised partials for this split ----
auto split_slot = [&](int h) -> size_t {
size_t bh = (size_t)batch * p.q_head + h;
return bh * p.num_splits + split;
};
#pragma unroll
for (int dn8 = 0; dn8 < 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;
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;
op[d] = Oacc[dn8][2];
op[d + 1] = Oacc[dn8][3];
}
}
if (tid4 == 0) {
int r0 = gid, r1 = gid + 8;
if (r0 < G) {
int h = q_head0 + r0;
float* mp = p.ml_part + split_slot(h) * 2;
mp[0] = m0; mp[1] = l0;
}
if (r1 < G) {
int h = q_head0 + r1;
float* mp = p.ml_part + split_slot(h) * 2;
mp[0] = m1; mp[1] = l1;
}
}
}
+177
View File
@@ -0,0 +1,177 @@
#pragma once
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include "attn_common.h"
using bf16 = __nv_bfloat16;
inline int compute_num_splits(int base_blocks, int tiles_total) {
int sm_count = 0;
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
return std::max(1, std::min(n, std::min(tiles_total, 32)));
}
// Dispatch head_dim: shared macro — avoids C++20 lambda template syntax.
// Usage: DISPATCH_HEAD_DIM(hd, fn, arg)
// Expands to: fn<32>(arg); fn<64>(arg); etc.
#define DISPATCH_HEAD_DIM(hd, fn, arg) \
switch (hd) { \
case 32: fn<32>(arg); break; \
case 64: fn<64>(arg); break; \
case 128: fn<128>(arg); break; \
case 256: fn<256>(arg); break; \
default: \
TORCH_CHECK(false, "unsupported head_dim ", hd, \
" (supported: 32, 64, 128, 256)"); \
}
template<typename P>
inline void alloc_split_partials(P& p) {
auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
auto o_part = torch::empty({p.batch, p.q_head, p.num_splits, p.head_dim}, fopt);
auto ml_part = torch::empty({p.batch, p.q_head, p.num_splits, 2}, fopt);
p.o_part = (float*)o_part.data_ptr();
p.ml_part = (float*)ml_part.data_ptr();
}
// ---- Shared Q-dims + strides extraction ----
template <typename P>
inline void extract_q_dims_and_strides(torch::Tensor& q, int64_t layout, P& p) {
if (layout == 1) q = q.transpose(1, 2);
p.batch = (int)q.size(0);
p.q_head = (int)q.size(1);
p.q_len = (int)q.size(2);
p.head_dim = (int)q.size(3);
p.q_stride_b = (int)q.stride(0);
p.q_stride_h = (int)q.stride(1);
p.q_stride_l = (int)q.stride(2);
p.q_stride_d = (int)q.stride(3);
}
// ---- Shared mask packing ----
template <typename P>
inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
if (p.use_mask) {
auto m = mask.value();
TORCH_CHECK(m.is_cuda(), "mask must be on CUDA");
TORCH_CHECK(m.dtype() == torch::kBool, "mask must be bool");
TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch");
TORCH_CHECK(m.size(m.dim() - 1) == p.kv_len, "mask kv_len mismatch");
if (m.dim() == 2) {
p.mask_b_stride = (int)m.stride(0);
p.mask_q_stride = 0;
} else if (m.dim() == 3) {
TORCH_CHECK(m.size(1) == p.q_len, "mask q_len mismatch");
p.mask_b_stride = (int)m.stride(0);
p.mask_q_stride = (int)m.stride(1);
} else {
TORCH_CHECK(false, "mask must be 2D [batch, kv_len] or 3D [batch, q_len, kv_len]");
}
p.mask = m.data_ptr<bool>();
} else {
p.mask = nullptr;
p.mask_b_stride = 0;
p.mask_q_stride = 0;
}
}
// ---- attn_pack_params (contiguous KV) ----
template<typename T>
inline void attn_pack_params(
torch::Tensor q,
torch::Tensor k,
torch::Tensor v,
c10::optional<torch::Tensor> mask,
int64_t causal_offset,
double scale,
int64_t layout,
AttentionParams<T>& p
) {
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
TORCH_CHECK(q.is_cuda() && k.is_cuda() && v.is_cuda());
TORCH_CHECK(q.dtype() == torch::kBFloat16);
TORCH_CHECK(k.dtype() == torch::kBFloat16);
TORCH_CHECK(v.dtype() == torch::kBFloat16);
TORCH_CHECK(k.sizes() == v.sizes(), "K and V must have identical shapes");
TORCH_CHECK(q.dim() == 4 && k.dim() == 4, "Q/K/V must be 4D");
extract_q_dims_and_strides(q, layout, p);
if (layout == 1) k = k.transpose(1, 2), v = v.transpose(1, 2);
p.kv_head = (int)k.size(1);
p.kv_len = (int)k.size(2);
TORCH_CHECK(k.size(3) == p.head_dim, "K/V head_dim must match Q");
p.kv_stride_b = (int)k.stride(0);
p.kv_stride_h = (int)k.stride(1);
p.kv_stride_l = (int)k.stride(2);
p.kv_stride_d = (int)k.stride(3);
p.causal_offset = (int)causal_offset;
p.use_mask = mask.has_value() ? 1 : 0;
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
p.q = (const T*)q.data_ptr();
p.k = (const T*)k.data_ptr();
p.v = (const T*)v.data_ptr();
p.o = nullptr;
p.o_part = nullptr;
p.ml_part = nullptr;
pack_mask(mask, p);
}
// ---- attn_pack_paged_params ----
template<typename T>
inline void attn_pack_paged_params(
torch::Tensor q,
torch::Tensor page_table,
torch::Tensor k_cache,
torch::Tensor v_cache,
int64_t page_size,
int64_t kv_len,
c10::optional<torch::Tensor> mask,
int64_t causal_offset,
double scale,
int64_t layout,
PagedAttentionParams<T>& p
) {
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
TORCH_CHECK(q.is_cuda() && page_table.is_cuda() && k_cache.is_cuda() && v_cache.is_cuda());
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16");
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16");
TORCH_CHECK(page_table.dtype() == torch::kLong, "page_table must be int64");
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must have identical shapes");
extract_q_dims_and_strides(q, layout, p);
p.kv_head = (int)k_cache.size(2);
p.kv_len = (int)kv_len;
p.page_size = (int)page_size;
p.max_pages = (int)page_table.size(1);
TORCH_CHECK(q.size(2) == 1, "Q seq_len must be 1 (decode)");
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
TORCH_CHECK(k_cache.size(1) == page_size,
"k_cache dim 1 must equal page_size, got ",
k_cache.size(1), " vs ", page_size);
p.causal_offset = (int)causal_offset;
p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0;
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
p.page_table = page_table.data_ptr<int64_t>();
p.k_cache = (const T*)k_cache.data_ptr();
p.v_cache = (const T*)v_cache.data_ptr();
p.q = (const T*)q.data_ptr();
p.o = nullptr;
p.o_part = nullptr;
p.ml_part = nullptr;
pack_mask(mask, p);
}
+293
View File
@@ -0,0 +1,293 @@
#pragma once
#include <cfloat>
#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.
// 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(
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};"
: "=f"(d[0]), "=f"(d[1]), "=f"(d[2]), "=f"(d[3])
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]),
"f"(c[0]), "f"(c[1]), "f"(c[2]), "f"(c[3]));
}
// read two adjacent bf16 from smem as one packed .b32 (elem0 low, elem1 high)
__device__ __forceinline__ unsigned ld2(const bf16* p) {
return *reinterpret_cast<const unsigned*>(p);
}
// pack two floats into one bf16x2 as .b32
__device__ __forceinline__ unsigned pk2(float a, float b) {
__nv_bfloat162 v = __floats2bfloat162_rn(a, b);
return *reinterpret_cast<unsigned*>(&v);
}
// pack two (non-contiguous) bf16 into one .b32
__device__ __forceinline__ unsigned pkb(bf16 a, bf16 b) {
__nv_bfloat162 v;
v.x = a;
v.y = b;
return *reinterpret_cast<unsigned*>(&v);
}
// 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.
__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];"
: "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3])
: "r"(a));
}
__device__ __forceinline__ void ldmatrix_x2(unsigned* r, const bf16* p) {
unsigned a = __cvta_generic_to_shared(p);
asm volatile("ldmatrix.sync.aligned.m8n8.x2.shared.b16 {%0,%1}, [%2];"
: "=r"(r[0]), "=r"(r[1])
: "r"(a));
}
__device__ __forceinline__ void ldmatrix_x2_trans(unsigned* r, const bf16* p) {
unsigned a = __cvta_generic_to_shared(p);
asm volatile("ldmatrix.sync.aligned.m8n8.x2.trans.shared.b16 {%0,%1}, [%2];"
: "=r"(r[0]), "=r"(r[1])
: "r"(a));
}
// 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+.
__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.
__device__ __forceinline__ void cp_async_16_pred(bf16* smem_ptr,
const void* gmem_ptr,
bool pred) {
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
int src_size = pred ? 16 : 0;
asm volatile("cp.async.ca.shared.global [%0], [%1], 16, %2;"
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
}
__device__ __forceinline__ void cp_async_commit() {
asm volatile("cp.async.commit_group;");
}
__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));
}
// ---------------------------------------------------------------------------
// Q-load: load query rows directly from global memory into mma A-operand
// register layout. One call replaces ~15 duplicated lines in each MMA kernel.
// stride_row is p.q_stride_h for decode (q_len=1, G heads) or
// p.q_stride_l for prefill (multi-q rows).
// ---------------------------------------------------------------------------
template <int KD>
__device__ inline void load_q_mma_frags(
const bf16* __restrict__ q,
int stride_row,
int stride_d,
int qra, int qrb,
bool va, bool vb,
int tid4,
unsigned Qa[KD][4])
{
#pragma unroll
for (int kt = 0; kt < KD; kt++) {
int c = kt * 16 + tid4 * 2;
const unsigned* pau = reinterpret_cast<const unsigned*>(
&q[qra * stride_row + c * stride_d]);
const unsigned* pbu = reinterpret_cast<const unsigned*>(
&q[qrb * stride_row + c * stride_d]);
Qa[kt][0] = va ? pau[0] : 0u;
Qa[kt][1] = vb ? pbu[0] : 0u;
Qa[kt][2] = va ? pau[4] : 0u;
Qa[kt][3] = vb ? pbu[4] : 0u;
}
}
// ---------------------------------------------------------------------------
// Shared MMA compute functions — used by both decode and prefill MMA kernels.
// Extracted because S=Q@K^T, online softmax, and P@V are structurally identical
// between the two kernels; only the per-row causal/mask bounds differ.
// ---------------------------------------------------------------------------
// S = Q @ K^T (Qa pre-loaded by the caller; scale applied post-mma in the
// caller to avoid bf16 precision loss).
// LD and SWIZ_MASK are constexpr in the calling kernel — passing them as
// runtime ints lets the compiler fold them while keeping the signature clean.
template <int KD, int NC8>
__device__ inline void mma_compute_scores(
const unsigned Qa[KD][4],
const bf16* __restrict__ sK,
int LD,
int SWIZ_MASK,
int lane,
float Sacc[NC8][4])
{
#pragma unroll
for (int n8 = 0; n8 < 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++) {
unsigned b[2];
ldmatrix_x2(b, &sK[krow_l * LD + swiz_col(kt * 16 + kcol_h, krow_l, 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>
__device__ inline void mma_softmax_tile(
int kv0,
int maxc0,
int maxc1,
int qrow0,
int qrow1,
int mask_b_stride,
int mask_q_stride,
int mask_batch,
const bool* __restrict__ mask,
bool has_mask,
float Sacc[NC8][4],
float Oacc[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++) {
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]);
float s0 = b0 ? -FLT_MAX : Sacc[n8][0];
float s1 = b1 ? -FLT_MAX : Sacc[n8][1];
float s2 = b2 ? -FLT_MAX : Sacc[n8][2];
float s3 = b3 ? -FLT_MAX : Sacc[n8][3];
Sacc[n8][0] = s0; Sacc[n8][1] = s1;
Sacc[n8][2] = s2; Sacc[n8][3] = s3;
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++) {
float p0 = pn0 * __expf(Sacc[n8][0] - nm0);
float p1 = pn0 * __expf(Sacc[n8][1] - nm0);
float p2 = pn1 * __expf(Sacc[n8][2] - nm1);
float p3 = pn1 * __expf(Sacc[n8][3] - nm1);
Sacc[n8][0] = p0; Sacc[n8][1] = p1;
Sacc[n8][2] = p2; Sacc[n8][3] = p3;
rsum0 += p0 + p1;
rsum1 += p2 + p3;
}
rsum0 += __shfl_xor_sync(0xFFFFFFFF, rsum0, 1);
rsum0 += __shfl_xor_sync(0xFFFFFFFF, rsum0, 2);
rsum1 += __shfl_xor_sync(0xFFFFFFFF, rsum1, 1);
rsum1 += __shfl_xor_sync(0xFFFFFFFF, rsum1, 2);
l0 = l0 * corr0 + rsum0;
l1 = l1 * corr1 + rsum1;
m0 = nm0; m1 = nm1;
#pragma unroll
for (int j = 0; j < 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>
__device__ inline void mma_pv_accumulate(
float Sacc[][4],
const bf16* __restrict__ sV,
int LD, int SWIZ_MASK, int lane,
float Oacc[DN8][4])
{
#pragma unroll
for (int kt2 = 0; kt2 < 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]);
Pa[2] = pk2(Sacc[kt2 * 2 + 1][0], Sacc[kt2 * 2 + 1][1]);
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++) {
unsigned b[2];
ldmatrix_x2_trans(b, &sV[vrow_l * LD + swiz_col(dn8 * 8, vrow_l, SWIZ_MASK)]);
mma16816(Oacc[dn8], Pa, b, Oacc[dn8]);
}
}
}
+82
View File
@@ -0,0 +1,82 @@
#include "attn_paged_decode_split_kv.cuh"
#ifndef ASTRAI_NO_MMA
#include "attn_paged_decode_split_kv_mma.cuh"
#endif
#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,
torch::Tensor k_cache,
torch::Tensor v_cache,
int64_t page_size,
int64_t kv_len,
c10::optional<torch::Tensor> mask,
int64_t causal_offset,
double scale,
int64_t layout
) {
PagedAttentionParams<bf16> p;
attn_pack_paged_params(q, page_table, k_cache, v_cache,
page_size, kv_len, mask, causal_offset, scale, layout, p);
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
p.o = (bf16*)O_view.data_ptr();
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p);
return O;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("attn_paged_decode", &attn_paged_decode,
py::arg("q"),
py::arg("page_table"),
py::arg("k_cache"),
py::arg("v_cache"),
py::arg("page_size"),
py::arg("kv_len"),
py::arg("mask") = py::none(),
py::arg("causal_offset") = -1,
py::arg("scale") = 0.0,
py::arg("layout") = 0,
"Paged GQA decode — split-KV with direct page-table access.");
}
+147
View File
@@ -0,0 +1,147 @@
#pragma once
#include <cuda_bf16.h>
#include <float.h>
#include "attn_common.h"
using bf16 = __nv_bfloat16;
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.
__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;
int split = blockIdx.z;
int group_size = blockDim.y;
int q_head = kv_head * group_size + threadIdx.y;
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;
#pragma unroll
for (int i = 0; i < hd_per_thread; i++)
q_reg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
extern __shared__ __align__(16) bf16 k_smem[];
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
int chunks_per_split = (chunks_total + p.num_splits - 1) / p.num_splits;
int ch_begin = split * chunks_per_split;
int ch_end = min(chunks_total, ch_begin + chunks_per_split);
const int mask_base = batch * p.mask_b_stride;
for (int ci = ch_begin; ci < ch_end; ci++) {
int chunk_start = ci * PDC_CHUNK;
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) {
int s = i / p.head_dim;
int d_dim = i % p.head_dim;
int pos = chunk_start + s;
int logical_page = pos / p.page_size;
int page_offset = pos % p.page_size;
int phys_page = p.page_table[batch * p.max_pages + logical_page];
if (phys_page >= 0) {
int64_t off = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim
+ (int64_t)page_offset * p.kv_head * p.head_dim
+ (int64_t)kv_head * p.head_dim
+ d_dim;
k_smem[i] = p.k_cache[off];
} else {
k_smem[i] = __float2bfloat16(0.0f);
}
}
__syncthreads();
for (int s = 0; s < this_chunk; s++) {
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;
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;
float new_m = fmaxf(m, partial);
float alpha = expf(m - new_m);
float beta = expf(partial - new_m);
d = d * alpha + beta;
int pos = chunk_start + s;
int logical_page = pos / p.page_size;
int page_offset = pos % p.page_size;
int phys_page = p.page_table[batch * p.max_pages + logical_page];
if (phys_page >= 0) {
int64_t v_base = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim
+ (int64_t)page_offset * p.kv_head * p.head_dim
+ (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;
} else {
#pragma unroll
for (int i = 0; i < hd_per_thread; i++)
acc_reg[i] = acc_reg[i] * alpha + 0.0f * beta;
}
m = new_m;
}
__syncthreads();
}
size_t bh = (size_t)batch * p.q_head + q_head;
size_t slot = bh * p.num_splits + split;
int d0 = lane * hd_per_thread;
#pragma unroll
for (int i = 0; i < hd_per_thread; i++)
p.o_part[slot * p.head_dim + (d0 + i)] = acc_reg[i];
if (lane == 0) {
p.ml_part[slot * 2] = m;
p.ml_part[slot * 2 + 1] = d;
}
}
__global__ void paged_attn_decode_combine_kernel(PagedAttentionParams<bf16> p) {
int bh = blockIdx.x;
int d = threadIdx.x;
if (d >= p.head_dim) return;
int batch = bh / p.q_head;
int q_head = bh % p.q_head;
size_t split_base = (size_t)bh * p.num_splits;
const float* mlp = p.ml_part + split_base * 2;
const float* op = p.o_part + split_base * p.head_dim;
float m = -FLT_MAX, l = 0.0f, acc = 0.0f;
for (int s = 0; s < p.num_splits; s++) {
float mi = mlp[s * 2];
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;
m = nm;
}
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
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);
}
@@ -0,0 +1,170 @@
#pragma once
#include <cfloat>
#include <cuda_bf16.h>
#include "attn_common.h"
#include "attn_mma_utils.cuh"
using bf16 = __nv_bfloat16;
// 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>
__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 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];
// ---- 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);
float Oacc[DN8][4];
#pragma unroll
for (int j = 0; j < DN8; j++)
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
const int tiles_total = (p.kv_len + BC - 1) / BC;
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
const int ti_begin = split * tiles_per_split;
const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
const int has_mask = p.use_mask && p.mask;
// 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;
// ---- Load tile lambda: predicated cp.async, 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 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;
int kc = kv0 + r;
bool valid = (kc < p.kv_len) && page_valid;
int page_off = kc % p.page_size;
int64_t gmem_base = (int64_t)phys_page * page_stride
+ (int64_t)page_off * pos_stride
+ head_off;
int off = r * LD + swiz_col(d, r, SWIZ_MASK);
cp_async_16_pred(&dK[off], &p.k_cache[gmem_base + d], valid);
cp_async_16_pred(&dV[off], &p.v_cache[gmem_base + d], valid);
}
cp_async_commit();
};
// ---- Prologue: issue first tile load ----
if (ti_begin < ti_end) {
load_tile(ti_begin, 0);
}
for (int ti = ti_begin; ti < ti_end; ti++) {
constexpr int BUF_MASK = (STAGES > 1) ? (STAGES - 1) : 0;
int buf = (ti - ti_begin) & BUF_MASK;
cp_async_wait_group<0>();
__syncwarp();
if constexpr (STAGES > 1) {
if (ti + 1 < ti_end)
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
}
const bf16* bK = sK + buf * BC * LD;
const bf16* bV = sV + buf * BC * LD;
int kv0 = ti * BC;
float Sacc[NC8][4];
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
#pragma unroll
for (int n8 = 0; n8 < NC8; n8++)
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
// 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);
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc);
__syncwarp();
if constexpr (STAGES == 1) {
if (ti + 1 < ti_end)
load_tile(ti + 1, 0);
}
}
// ---- write UN-normalised partials for this split ----
auto split_slot = [&](int h) -> size_t {
size_t bh = (size_t)batch * p.q_head + h;
return bh * p.num_splits + split;
};
#pragma unroll
for (int dn8 = 0; dn8 < 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;
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;
op[d] = Oacc[dn8][2];
op[d + 1] = Oacc[dn8][3];
}
}
if (tid4 == 0) {
int r0 = gid, r1 = gid + 8;
if (r0 < G) {
int h = q_head0 + r0;
float* mp = p.ml_part + split_slot(h) * 2;
mp[0] = m0; mp[1] = l0;
}
if (r1 < G) {
int h = q_head0 + r1;
float* mp = p.ml_part + split_slot(h) * 2;
mp[0] = m1; mp[1] = l1;
}
}
}
+64
View File
@@ -0,0 +1,64 @@
#include "attn_prefill_split_q.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,
torch::Tensor v,
c10::optional<torch::Tensor> mask,
int64_t causal_offset,
double scale,
int64_t layout
) {
AttentionParams<bf16> p;
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16");
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
p.o = (bf16*)O_view.data_ptr();
DISPATCH_HEAD_DIM(p.head_dim, dispatch_prefill, p);
return O;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("attn_prefill", &attn_prefill,
py::arg("q"),
py::arg("k"),
py::arg("v"),
py::arg("mask") = py::none(),
py::arg("causal_offset") = -1,
py::arg("scale") = 0.0,
py::arg("layout") = 0,
"GQA prefill (tensor-core mma on sm_80+, scalar fallback)");
}
+152
View File
@@ -0,0 +1,152 @@
#pragma once
#include <cfloat>
#include <cuda_bf16.h>
#include "attn_common.h"
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.
template <int G>
__device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
#pragma unroll
for (int o = G / 2; o > 0; o >>= 1)
v += __shfl_xor_sync(mask, v, o);
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.
__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);
#pragma unroll
for (int j = 0; j < 4; j++) {
float2 f = __bfloat1622float2(h[j]);
o[2 * j] = f.x;
o[2 * j + 1] = f.y;
}
}
template <int HEAD_DIM, int G, int ROWS, int P_BC>
__global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
constexpr int DPT = HEAD_DIM / G;
int q_tile = blockIdx.x;
int q_head = blockIdx.y;
int batch = blockIdx.z;
int gpos = threadIdx.x; // 0..G-1 (which d-chunk)
int row = threadIdx.y; // 0..ROWS-1
int q_row = q_tile * ROWS + row;
int kv_head = q_head / (p.q_head / p.kv_head);
__shared__ __align__(16) bf16 sK[P_BC * HEAD_DIM];
__shared__ __align__(16) bf16 sV[P_BC * HEAD_DIM];
// Q: stride-based load [batch, q_head, q_len, head_dim]
float qreg[DPT];
if (q_row < p.q_len) {
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
+ q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
#pragma unroll
for (int i = 0; i < DPT; i++)
qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]) * p.scale;
}
float m = -FLT_MAX, l = 0.0f;
float acc[DPT];
#pragma unroll
for (int i = 0; i < DPT; i++)
acc[i] = 0.0f;
// KV: stride-based base
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
int mask_batch_base = batch * p.mask_b_stride;
int tiles = (p.kv_len + P_BC - 1) / P_BC;
int tt = G * ROWS;
int lid = row * G + gpos;
// 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)));
for (int ti = 0; ti < tiles; ti++) {
int kv0 = ti * P_BC;
int tlen = min(P_BC, p.kv_len - kv0);
// Load K/V into shared memory from strided global
for (int i = lid; i < tlen * HEAD_DIM; i += tt) {
int s = i / HEAD_DIM;
int d_dim = i % HEAD_DIM;
int kv_idx = kv0 + s;
int g_off = kv_base + kv_idx * p.kv_stride_l + d_dim * p.kv_stride_d;
sK[i] = p.k[g_off];
sV[i] = p.v[g_off];
}
__syncthreads();
int lim = tlen;
if (p.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;
}
int mask_row_base = mask_batch_base + q_row * p.mask_q_stride;
for (int s = 0; s < lim; s++) {
const bf16* kr = sK + s * HEAD_DIM + gpos * DPT;
float part = 0.0f;
#pragma unroll
for (int i = 0; i < DPT; i += 8) {
float k8[8];
ld8(kr + i, k8);
#pragma unroll
for (int j = 0; j < 8; j++)
part = fmaf(qreg[i + j], k8[j], part);
}
float dot = group_reduce_sum<G>(part, gmask);
int kv_idx = kv0 + s;
if (p.use_mask && p.mask && !p.mask[mask_row_base + kv_idx])
dot = -FLT_MAX;
float nm = fmaxf(m, dot);
float al = __expf(m - nm);
float be = __expf(dot - nm);
l = l * al + be;
const bf16* vr = sV + s * HEAD_DIM + gpos * DPT;
#pragma unroll
for (int i = 0; i < DPT; i += 8) {
float v8[8];
ld8(vr + i, v8);
#pragma unroll
for (int j = 0; j < 8; j++)
acc[i + j] = fmaf(v8[j], be, acc[i + j] * al);
}
m = nm;
}
__syncthreads();
}
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;
#pragma unroll
for (int i = 0; i < DPT; i++)
p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * rl);
}
}
+196
View File
@@ -0,0 +1,196 @@
#pragma once
#include <cfloat>
#include <cuda_bf16.h>
#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.
//
// 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.
//
// 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>
__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 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;
// ---- 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];
// 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);
float Oacc[DN8][4];
#pragma unroll
for (int j = 0; j < DN8; j++)
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
// KV: stride-based base
const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
const int tiles = (p.kv_len + BC - 1) / BC;
const int qr0 = qrow0 + gid; // row for c0/c1
const int qr1 = qrow0 + gid + 8; // row for c2/c3
// Causal tile-skip bounds (no-op when causal_offset < 0)
const int use_skip = (p.causal_offset >= 0) ? 1 : 0;
const int max_kv = qrow0 + BR - 1 + p.causal_offset;
const int block_max_kv =
blockIdx.x * WARPS * BR + WARPS * BR - 1 + p.causal_offset;
const int has_mask = p.use_mask && p.mask;
// 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 (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 kc = kv0 + r;
bool valid = kc < p.kv_len;
int off = r * LD + swiz_col(d, r, 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);
}
cp_async_commit();
};
// ---- Prologue: issue first tile load ----
load_tile(0, 0);
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.
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;
// Warp-level causal skip
if (!use_skip || 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);
// 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;
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);
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc);
} // if active (warp-level causal skip)
}
// ---- write output ---- (packed bf16x2 stores: one 32-bit STG per pair,
// halves store count and removes the uncoalesced scalar-store penalty)
float rl0 = (l0 > 1e-20f) ? (1.0f / l0) : 0.0f;
float rl1 = (l1 > 1e-20f) ? (1.0f / l1) : 0.0f;
// O: stride-based write
const int o_base = batch * p.q_stride_b + q_head * p.q_stride_h;
#pragma unroll
for (int dn8 = 0; dn8 < DN8; dn8++) {
int d = dn8 * 8 + 2 * tid4;
if (qr0 < p.q_len) {
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][0] * rl0,
Oacc[dn8][1] * rl0);
*reinterpret_cast<__nv_bfloat162*>(&p.o[o_base + qr0 * 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;
}
}
}
+201
View File
@@ -0,0 +1,201 @@
/*
Pure-C test:
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
// 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).
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;
}
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); });
}
// 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, 1024, 128},
{1, 32, 4, 2048, 128},
{1, 32, 4, 4096, 128},
{16, 32, 4, 2048, 128},
{32, 32, 4, 1024, 128},
};
const int WARMUP = 10, ITERS = 100;
printf("\n===== DECODE BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
print_bench_header();
for (int ci = 0; ci < 6; ci++) {
int B = cfgs[ci][0], Hq = cfgs[ci][1], Hk = cfgs[ci][2];
int sl = cfgs[ci][3], D = cfgs[ci][4];
size_t nQ = (size_t)B * Hq * D;
size_t nKV = (size_t)B * Hk * sl * D;
bf16 *dQ, *dK, *dV, *dO;
cudaMalloc(&dQ, nQ*2); cudaMalloc(&dK, nKV*2);
cudaMalloc(&dV, nKV*2); cudaMalloc(&dO, nQ*2);
size_t big = nQ > nKV ? nQ : nKV; bf16* tmp = new bf16[big];
for (size_t i = 0; i < nQ; i++) tmp[i] = f2bf(randf());
cudaMemcpy(dQ, tmp, nQ*2, cudaMemcpyHostToDevice);
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
cudaMemcpy(dK, tmp, nKV*2, cudaMemcpyHostToDevice);
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
cudaMemcpy(dV, tmp, nKV*2, cudaMemcpyHostToDevice);
delete[] tmp;
AttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hk; p.q_len = 1; p.kv_len = sl;
p.head_dim = D; p.use_mask = 0; p.causal_offset = -1;
p.scale = 1.0f / sqrtf((float)D);
set_default_strides(p);
p.q = dQ; p.k = dK; p.v = dV; p.mask = nullptr; p.o = dO;
DecodeScratch sc;
cudaMalloc(&sc.o_part, (size_t)B*Hq*32*D*sizeof(float));
cudaMalloc(&sc.ml_part, (size_t)B*Hq*32*2*sizeof(float));
auto launch = [&]() { dispatch_decode(p, sc); };
double flops = 4.0 * B * Hq * (double)sl * D;
double bytes = 2.0 * (2.0 * nKV * sizeof(bf16));
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops, bytes);
char cfg[64];
snprintf(cfg, sizeof(cfg),
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
B, Hq, Hk, 1, sl, D, 0);
print_bench_row(cfg, r);
cudaFree(dQ); cudaFree(dK); cudaFree(dV); cudaFree(dO);
cudaFree(sc.o_part); cudaFree(sc.ml_part);
}
}
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},
};
int n_cfgs = sizeof(configs) / sizeof(configs[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);
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;
}
printf("All tests passed!\n");
bench();
return 0;
}
+332
View File
@@ -0,0 +1,332 @@
// Compile:
// nvcc -I csrc -arch=sm_89 -O3 --use_fast_math --ptxas-options=-O3 \
// --extra-device-vectorization csrc/tests/attn_paged_decode_test.cu \
// -o /tmp/test_paged && /tmp/test_paged
#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
// 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,
int page_size, int head_dim,
bf16* h_k, bf16* h_v)
{
int max_pages = (kv_len + page_size - 1) / page_size;
size_t page_stride = (size_t)page_size * Hkv * head_dim;
for (int b = 0; b < B; b++) {
for (int pos = 0; pos < kv_len; pos++) {
int log_pg = pos / page_size;
int pg_off = pos % page_size;
int phys = (int)h_pt[b * max_pages + log_pg];
for (int h = 0; h < Hkv; h++) {
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;
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));
}
}
}
}
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);
fflush(stdout);
int max_pages = (kv_len + page_size - 1) / page_size;
int n_phys_pages = B * max_pages;
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_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);
cudaMalloc(&d_op, sz_op);
cudaMalloc(&d_ml, sz_ml);
srand(seed);
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
bf16* h_q = (bf16*)malloc(sz_q);
for (int i = 0; i < B * Hq * HEAD_DIM; i++)
h_q[i] = __float2bfloat16(rnd());
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
bf16* h_k_pool = (bf16*)malloc(sz_kv);
bf16* h_v_pool = (bf16*)malloc(sz_kv);
size_t ps = (size_t)page_size * Hkv * HEAD_DIM;
for (int pg = 0; pg < n_phys_pages; pg++) {
for (int off = 0; off < page_size; off++) {
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;
h_k_pool[idx] = __float2bfloat16(v);
h_v_pool[idx] = __float2bfloat16(v * 0.3f);
}
}
}
}
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_pt = (int64_t*)malloc(sz_pt);
int next_pg = 0;
for (int b = 0; b < B; b++)
for (int p = 0; p < max_pages; p++)
h_pt[b * max_pages + p] = next_pg++;
cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice);
bf16* h_k_cont = (bf16*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(bf16));
bf16* h_v_cont = (bf16*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(bf16));
gather_kv_cpu(h_k_pool, h_v_pool, h_pt, B, Hkv, kv_len, page_size, HEAD_DIM, h_k_cont, h_v_cont);
float* h_q_f = (float*)malloc((size_t)B * Hq * HEAD_DIM * sizeof(float));
float* h_k_f = (float*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(float));
float* h_v_f = (float*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(float));
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
for (int i = 0; i < B * kv_len * Hkv * HEAD_DIM; i++) {
h_k_f[i] = bf2f(h_k_cont[i]);
h_v_f[i] = bf2f(h_v_cont[i]);
}
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);
float scale_val = 1.0f / sqrtf((float)HEAD_DIM);
PagedAttentionParams<bf16, float> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.q_len = 1;
p.kv_len = kv_len; p.head_dim = HEAD_DIM;
p.use_mask = 0; p.causal_offset = -1;
set_default_paged_strides(p);
p.num_splits = 1; p.scale = scale_val;
p.page_size = page_size; p.max_pages = max_pages;
p.page_table = d_pt;
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);
cudaDeviceSynchronize();
bf16* h_o_bf16 = (bf16*)malloc(sz_o);
cudaMemcpy(h_o_bf16, d_o_paged, sz_o, cudaMemcpyDeviceToHost);
float* h_o_paged = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
for (int i = 0; i < B * Hq * HEAD_DIM; i++)
h_o_paged[i] = __bfloat162float(h_o_bf16[i]);
float max_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; }
}
bool pass = max_err < 0.02f;
if (pass) {
printf("PASS (max_abs_err=%.4e)\n", max_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(" ref[0..7]:");
for (int i = 0; i < 8 && i < HEAD_DIM; i++)
printf(" %.4f", h_o_ref[i]);
printf("\n got[0..7]:");
for (int i = 0; i < 8 && i < HEAD_DIM; i++)
printf(" %.4f", h_o_paged[i]);
printf("\n");
}
free(h_q); free(h_k_pool); free(h_v_pool); free(h_pt);
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_k_pool); cudaFree(d_v_pool); cudaFree(d_pt);
cudaFree(d_op); cudaFree(d_ml);
return pass ? 0 : 1;
}
struct TestCase {
int head_dim;
int B, Hq, Hkv, kv_len, page_size, 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
};
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);
});
return matched ? r : 1;
}
// 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;
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);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t* d_pt;
float *d_op, *d_ml;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
cudaMalloc(&d_pt, sz_pt);
cudaMalloc(&d_op, sz_op); cudaMalloc(&d_ml, sz_ml);
bf16* tmp = (bf16*)malloc(sz_kv > sz_q ? sz_kv : sz_q);
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) tmp[i] = f2bf(randf());
cudaMemcpy(d_q, tmp, sz_q, cudaMemcpyHostToDevice);
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) tmp[i] = f2bf(randf());
cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_pt = (int64_t*)malloc(sz_pt);
int next_pg = 0;
for (int b = 0; b < B; b++)
for (int p = 0; p < max_pages; p++)
h_pt[b * max_pages + p] = next_pg++;
cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice);
free(h_pt);
float scale_val = 1.0f / sqrtf((float)HEAD_DIM);
PagedAttentionParams<bf16, float> pa;
pa.batch = B; pa.q_head = Hq; pa.kv_head = Hkv; pa.q_len = 1;
pa.kv_len = kv_len; pa.head_dim = HEAD_DIM;
pa.use_mask = 0; pa.causal_offset = -1;
set_default_paged_strides(pa);
pa.num_splits = 1; pa.scale = scale_val;
pa.page_size = page_size; pa.max_pages = max_pages;
pa.page_table = d_pt;
pa.k_cache = d_k_pool; pa.v_cache = d_v_pool;
pa.q = d_q; pa.mask = nullptr; pa.o = d_o;
pa.o_part = d_op; pa.ml_part = d_ml;
const int WARMUP = 10, ITERS = 100;
auto launch = [&]() { launch_paged_decode<HEAD_DIM>(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));
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops, bytes);
char cfg[64];
snprintf(cfg, sizeof(cfg),
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d page=%3d",
B, Hq, Hkv, 1, kv_len, HEAD_DIM, page_size);
print_bench_row(cfg, r);
free(tmp);
cudaFree(d_q); cudaFree(d_o);
cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_pt);
cudaFree(d_op); cudaFree(d_ml);
}
static void bench() {
printf("\n===== PAGED DECODE BENCH =====\n");
print_bench_header();
bench_config<128>(1, 32, 4, 512, 128);
bench_config<128>(1, 32, 4, 1024, 128);
bench_config<128>(1, 32, 4, 2048, 128);
bench_config<128>(1, 32, 4, 4096, 128);
bench_config<128>(16, 32, 4, 2048, 128);
bench_config<128>(32, 32, 4, 1024, 128);
}
int main() {
int n = sizeof(TESTS) / sizeof(TESTS[0]);
int fail = 0;
printf("=== Paged Decode vs CPU reference (%d cases) ===\n\n", n);
for (int i = 0; i < n; i++) {
fail += dispatch_test(TESTS[i]);
if (fail) break;
}
if (fail) {
printf("\nFAILED (%d/%d tests failed)\n", fail, n);
return fail;
}
printf("\nAll %d tests passed!\n", n);
bench();
return 0;
}
+178
View File
@@ -0,0 +1,178 @@
/*
Pure-C test:
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);
}
}
// 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},
{1,32,4,1024,1024,128,0},
{1,32,4,2048,2048,128,0},
{1,32,4,2048,2048,128,1},
{4,32,4,2048,2048,128,1},
{1,32,4,4096,4096,128,1},
};
int n = sizeof(cfgs)/sizeof(cfgs[0]);
const int WARMUP = 10, ITERS = 50;
printf("\n===== PREFILL BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
printf("%-46s | %10s | %10s | %10s\n",
"config", "latency", "bandwidth", "throughput");
printf("---------------------------------------------------------------"
"----------------------------\n");
for (int ci = 0; ci < n; ci++) {
int B=cfgs[ci][0], Hq=cfgs[ci][1], Hk=cfgs[ci][2];
int ql=cfgs[ci][3], kl=cfgs[ci][4], D=cfgs[ci][5], causal=cfgs[ci][6];
size_t nQ=(size_t)B*Hq*ql*D, nKV=(size_t)B*Hk*kl*D;
bf16 *dQ,*dK,*dV,*dO,*tmp;
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
size_t big = nQ>nKV?nQ:nKV; tmp=new bf16[big];
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(randf());
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
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;
for (int i=0;i<WARMUP;i++) dispatch_prefill(p);
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);
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;
char cfg[64];
snprintf(cfg, sizeof(cfg),
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
B,Hq,Hk,ql,kl,D,causal);
printf("%-46s | %7.4f ms | %7.1f GB/s | %6.2f TFLOP/s\n",
cfg, ms, gbps, tflops);
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
delete[]tmp; cudaEventDestroy(s); cudaEventDestroy(e);
}
}
int main() {
const int configs[][7] = {
{1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal
{1,32,4,512,512,128,0}, // standard
{1,32,4,128,256,128,0}, // medium
{1,4,2,256,256,128,1}, // causal
};
int n_configs = sizeof(configs) / sizeof(configs[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);
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;
}
printf("All tests passed!\n");
bench();
return 0;
}
+181
View File
@@ -0,0 +1,181 @@
#pragma once
#include <cstdio>
#include <cstdlib>
#include <cmath>
#include <chrono>
#include <cuda_bf16.h>
using bf16 = __nv_bfloat16;
inline bf16 f2bf(float x) { return __float2bfloat16(x); }
inline float bf2f(bf16 x) { return __bfloat162float(x); }
inline float randf() { return (float)rand() / (float)RAND_MAX - 0.5f; }
inline double now_ms() {
using namespace std::chrono;
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); \
if (_e != cudaSuccess) { \
printf("CUDA error %s at %s:%d\n", cudaGetErrorString(_e), __FILE__, __LINE__); \
exit(1); \
} \
} while (0)
struct BenchResult {
float ms;
double gbps;
double tflops;
};
template <typename Fn>
BenchResult bench_kernel(Fn launch, int warmup, int iters,
double flops, double bytes) {
for (int i = 0; i < warmup; i++) launch();
cudaDeviceSynchronize();
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
printf("CUDA error before bench: %s\n", cudaGetErrorString(err));
return {0, 0, 0};
}
cudaEvent_t s, e;
cudaEventCreate(&s); cudaEventCreate(&e);
cudaEventRecord(s);
for (int i = 0; i < iters; i++) launch();
cudaEventRecord(e); cudaEventSynchronize(e);
float ms = 0; cudaEventElapsedTime(&ms, s, e); ms /= iters;
cudaEventDestroy(s); cudaEventDestroy(e);
return {ms, bytes / (ms * 1e-3) / 1e9, flops / (ms * 1e-3) / 1e12};
}
inline void print_bench_header() {
printf("%-46s | %10s | %10s | %10s\n",
"config", "latency", "bandwidth", "throughput");
printf("---------------------------------------------------------------"
"----------------------------\n");
}
inline void print_bench_row(const char* cfg, const BenchResult& r) {
printf("%-46s | %7.4f ms | %7.1f GB/s | %6.2f TFLOP/s\n",
cfg, r.ms, r.gbps, r.tflops);
}
template <int... Ds>
struct _HeadSwitch;
template <int D>
struct _HeadSwitch<D> {
template <typename Fn>
static void call(int hd, Fn&& fn) { if (hd == D) fn.template operator()<D>(); }
};
template <int D, int... Rest>
struct _HeadSwitch<D, Rest...> {
template <typename Fn>
static void call(int hd, Fn&& fn) {
if (hd == D) fn.template operator()<D>();
else _HeadSwitch<Rest...>::call(hd, fn);
}
};
// Default set: 32, 64, 128, 256
template <typename Fn>
void dispatch_by_head_dim(int head_dim, Fn&& fn) {
_HeadSwitch<32, 64, 128, 256>::call(head_dim, fn);
}
// Set default strides for contiguous b h l d layout on AttentionParams.
template<typename P>
inline void set_default_strides(P& p) {
p.q_stride_b = p.q_head * p.q_len * p.head_dim;
p.q_stride_h = p.q_len * p.head_dim;
p.q_stride_l = p.head_dim;
p.q_stride_d = 1;
p.kv_stride_b = p.kv_head * p.kv_len * p.head_dim;
p.kv_stride_h = p.kv_len * p.head_dim;
p.kv_stride_l = p.head_dim;
p.kv_stride_d = 1;
p.mask_b_stride = p.kv_len;
p.mask_q_stride = 0;
}
// Set default Q strides for contiguous b h l d layout on PagedAttentionParams.
template<typename P>
inline void set_default_paged_strides(P& p) {
p.q_stride_b = p.q_head * p.q_len * p.head_dim;
p.q_stride_h = p.q_len * p.head_dim;
p.q_stride_l = p.head_dim;
p.q_stride_d = 1;
p.mask_b_stride = p.kv_len;
p.mask_q_stride = 0;
}
// Generic CPU reference for multi-query / grouped-query attention.
// Tensor shapes (all float*):
// Q : [B, Hq, q_len, D]
// K : [B, Hk, kv_len, D]
// V : [B, Hk, kv_len, D]
// O : [B, Hq, q_len, D]
// mask: if q_len == 1, shape is [B, kv_len]; otherwise mask is not supported.
// causal_offset: -1 = non-causal; >=0 = absolute position of first Q token.
static void cpu_attention_ref(
const float* Q, const float* K, const float* V, const bool* mask,
float* O, int B, int Hq, int Hk, int q_len, int kv_len, int D,
int causal_offset
) {
float scale = 1.0f / sqrtf((float)D);
int n_rep = Hq / Hk;
for (int b = 0; b < B; b++) {
for (int h = 0; h < Hq; h++) {
int kv_h = h / n_rep;
for (int qi = 0; qi < q_len; qi++) {
float mv = -INFINITY, sv = 0.0f;
float accum[256] = {0.0f};
int lim = kv_len;
if (causal_offset >= 0) {
int c = qi + causal_offset + 1;
lim = (c < kv_len) ? c : kv_len;
}
for (int kj = 0; kj < lim; kj++) {
if (mask != nullptr && q_len == 1) {
if (!mask[b * kv_len + kj]) continue;
}
float dot = 0.0f;
size_t q_idx = ((size_t)b * Hq + h) * q_len + qi;
size_t kv_idx = ((size_t)b * Hk + kv_h) * kv_len + kj;
for (int d = 0; d < D; d++)
dot += Q[q_idx * D + d] * K[kv_idx * D + d];
dot *= scale;
float nm = fmaxf(mv, dot);
float a = expf(mv - nm);
float b_exp = expf(dot - nm);
sv = sv * a + b_exp;
for (int d = 0; d < D; d++)
accum[d] = accum[d] * a + V[kv_idx * D + d] * b_exp;
mv = nm;
}
float inv = 1.0f / sv;
size_t o_idx = ((size_t)b * Hq + h) * q_len + qi;
for (int d = 0; d < D; d++)
O[o_idx * D + d] = accum[d] * inv;
}
}
}
}
+3 -3
View File
@@ -9,8 +9,8 @@ readme = "README.md"
requires-python = ">=3.12" requires-python = ">=3.12"
dependencies = [ dependencies = [
"h5py==3.15.1", "h5py==3.15.1",
"numpy==2.3.2", "numpy==2.4.4",
"torch==2.7.1", "torch==2.11.0",
"tokenizers==0.21.4", "tokenizers==0.21.4",
"tqdm==4.67.1", "tqdm==4.67.1",
"safetensors==0.5.3", "safetensors==0.5.3",
@@ -37,7 +37,7 @@ dev = ["pytest==9.0.2", "ruff"]
where = ["."] where = ["."]
[tool.pip] [tool.pip]
extra-index-url = "https://download.pytorch.org/whl/cu126" extra-index-url = "https://download.pytorch.org/whl/cu128"
[tool.setuptools.dynamic] [tool.setuptools.dynamic]
version = { attr = "astrai.__version__" } version = { attr = "astrai.__version__" }
+1 -1
View File
@@ -5,7 +5,7 @@ from huggingface_hub import snapshot_download
PROJECT_ROOT = Path(__file__).resolve().parents[2] PROJECT_ROOT = Path(__file__).resolve().parents[2]
DEFAULT_LOCAL_DIR = Path(PROJECT_ROOT, "params") DEFAULT_LOCAL_DIR = Path(PROJECT_ROOT, "params")
DEFAULT_REPO_ID = "ViperEk/KHAOSZ" DEFAULT_REPO_ID = "ViperEkura/AstrAI-V1-instruct"
if __name__ == "__main__": if __name__ == "__main__":
parser = argparse.ArgumentParser( parser = argparse.ArgumentParser(
+2 -4
View File
@@ -26,11 +26,9 @@ def batch_generate():
prompts = [ prompts = [
tokenizer.apply_chat_template( tokenizer.apply_chat_template(
[ [{"role": "user", "content": q}],
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": q},
],
tokenize=False, tokenize=False,
add_generation_prompt=True,
) )
for q in inputs for q in inputs
] ]
+73 -16
View File
@@ -1,3 +1,4 @@
from argparse import ArgumentParser
from pathlib import Path from pathlib import Path
import torch import torch
@@ -7,15 +8,69 @@ from astrai.model import AutoModel
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
PROJECT_ROOT = Path(__file__).resolve().parents[2] PROJECT_ROOT = Path(__file__).resolve().parents[2]
PARAMETER_ROOT = Path(PROJECT_ROOT, "params")
def parse_args():
parser = ArgumentParser(description="Interactive streaming chat")
parser.add_argument(
"--model_path",
type=Path,
default=PROJECT_ROOT / "params",
help="Path to model weights (params/ or checkpoint/epoch_N_step_M/)",
)
parser.add_argument(
"--temperature",
type=float,
default=0.8,
help="Sampling temperature (default: 0.8)",
)
parser.add_argument(
"--top_p",
type=float,
default=0.95,
help="Top-p sampling threshold",
)
parser.add_argument(
"--top_k",
type=int,
default=50,
help="Top-k sampling threshold",
)
parser.add_argument(
"--max_tokens",
type=int,
default=2048,
help="Maximum tokens to generate",
)
parser.add_argument(
"--frequency_penalty",
type=float,
default=0.5,
help="Penalty per occurrence for repeated tokens (0.0 disables, "
"range -2.0~2.0, typical 0.3-1.0)",
)
parser.add_argument(
"--rep_window",
type=int,
default=64,
help="Number of recent prompt tokens to include in penalty history",
)
parser.add_argument(
"--system_prompt",
type=str,
default="",
help="Optional system prompt (default: empty, model not SFT-trained on system role)",
)
return parser.parse_args()
def chat(): def chat():
model = AutoModel.from_pretrained(PARAMETER_ROOT) args = parse_args()
tokenizer = AutoTokenizer.from_pretrained(PARAMETER_ROOT) model_path = args.model_path
model.to(device="cuda", dtype=torch.bfloat16)
messages = [{"role": "system", "content": "You are a helpful assistant."}] model = AutoModel.from_pretrained(model_path)
tokenizer = AutoTokenizer.from_pretrained(model_path)
model.to(device="cuda", dtype=torch.bfloat16)
engine = InferenceEngine(model=model, tokenizer=tokenizer) engine = InferenceEngine(model=model, tokenizer=tokenizer)
while True: while True:
@@ -23,27 +78,29 @@ def chat():
if query == "!exit": if query == "!exit":
break break
# Add user message msgs = []
messages.append({"role": "user", "content": query}) if args.system_prompt:
msgs.append({"role": "system", "content": args.system_prompt})
msgs.append({"role": "user", "content": query})
prompt = tokenizer.apply_chat_template(
msgs, tokenize=False, add_generation_prompt=True
)
# Generate response
full_response = "" full_response = ""
prompt = tokenizer.apply_chat_template(messages, tokenize=False)
for token in engine.generate( for token in engine.generate(
prompt=prompt, prompt=prompt,
stream=True, stream=True,
max_tokens=2048, max_tokens=args.max_tokens,
temperature=0.8, temperature=args.temperature,
top_p=0.95, top_p=args.top_p,
top_k=50, top_k=args.top_k,
frequency_penalty=args.frequency_penalty,
rep_window=args.rep_window,
): ):
print(token, end="", flush=True) print(token, end="", flush=True)
full_response += token full_response += token
print() print()
# Add assistant response to messages
messages.append({"role": "assistant", "content": full_response.strip()})
if __name__ == "__main__": if __name__ == "__main__":
+321
View File
@@ -0,0 +1,321 @@
"""SVD effective rank & weight statistics analysis for model checkpoints."""
import argparse
import json
from pathlib import Path
import safetensors.torch
import torch
def effective_rank_metrics(w: torch.Tensor) -> dict:
if w.ndim == 1:
return {"shape": tuple(w.shape), "is_1d": True}
w = w.float()
s = torch.linalg.svdvals(w)
s_sq = s**2
total = s_sq.sum()
cumsum = torch.cumsum(s_sq, dim=0) / total
min_dim = min(w.shape[0], w.shape[1])
er_90 = (cumsum < 0.90).sum().item() + 1
er_95 = (cumsum < 0.95).sum().item() + 1
er_99 = (cumsum < 0.99).sum().item() + 1
p = s_sq / total
p = p[p > 1e-30]
entropy = -(p * torch.log(p)).sum()
entropic_rank = torch.exp(entropy).item()
return {
"shape": tuple(w.shape),
"min_dim": min_dim,
"er_90": er_90,
"er_95": er_95,
"er_99": er_99,
"er_99_norm": er_99 / min_dim,
"er_95_norm": er_95 / min_dim,
"entropic_rank": entropic_rank,
"entropic_rank_norm": entropic_rank / min_dim,
"top1_ratio": s[0].item() / s.sum().item(),
"top5_ratio": s[:5].sum().item() / s.sum().item(),
"decay_ratio": s[-1].item() / s[0].item(),
"condition_number": s[0].item() / s[-1].item(),
"mean": w.mean().item(),
"std": w.std().item(),
"min": w.min().item(),
"max": w.max().item(),
}
def format_header(headers: list[str], widths: list[int]) -> str:
return "".join(h.ljust(w) for h, w in zip(headers, widths))
def format_row(values: list[str], widths: list[int]) -> str:
return "".join(v.ljust(w) for v, w in zip(values, widths))
def group_by_component(results: dict[str, dict]) -> dict[str, list[dict]]:
groups: dict[str, list[dict]] = {}
for key, r in results.items():
parts = key.split(".")
if parts[0] == "layers" and len(parts) >= 3:
sub = parts[2:]
if sub[0] == "attention":
comp = f"attn.{sub[1]}"
elif sub[0] == "mlp":
comp = f"mlp.{sub[1]}"
elif sub[0] == "input_norm":
comp = "input_norm"
elif sub[0] == "post_attention_norm":
comp = "post_attn_norm"
else:
comp = ".".join(sub)
else:
comp = key
groups.setdefault(comp, []).append(r)
return groups
def print_component_summary(results: dict[str, dict], title: str):
groups = group_by_component(results)
matrix_groups = {
k: [v for v in vs if not v.get("is_1d")]
for k, vs in groups.items()
if any(not v.get("is_1d") for v in vs)
}
widths = [20, 12, 12, 12, 12, 12]
print(f"\n{title}")
print(
format_header(
["Component", "N", "ER@99%", "EntRank%", "Top1 σ(%)", "Cond. Num"], widths
)
)
print("-" * sum(widths))
for name in sorted(matrix_groups.keys()):
items = matrix_groups[name]
n = len(items)
print(
format_row(
[
name,
str(n),
f"{sum(r['er_99_norm'] for r in items) / n:.4f}",
f"{sum(r['entropic_rank_norm'] for r in items) / n:.4f}",
f"{sum(r['top1_ratio'] for r in items) / n:.4f}",
f"{sum(r['condition_number'] for r in items) / n:.1f}",
],
widths,
)
)
all_er = [
r["er_99_norm"]
for vs in matrix_groups.values()
for r in vs
if not r.get("is_1d")
]
if all_er:
m = sum(all_er) / len(all_er)
print(f"\n Overall Mean ER@99: {m:.4f} ({m * 100:.1f}% of dimension)")
if m > 0.85:
print(" → HIGH utilization: model near capacity → need more params")
elif m > 0.5:
print(" → MODERATE utilization: some headroom left")
else:
print(" → LOW utilization: significant unused capacity")
def print_layer_grid(results: dict[str, dict]):
comps = [
"attn.q_proj",
"attn.k_proj",
"attn.v_proj",
"attn.o_proj",
"mlp.up",
"mlp.gate",
"mlp.down",
]
widths = [6] + [10] * len(comps)
metric = "er_99_norm"
print(f"\n--- Per-Layer Effective Rank (99% energy) ---")
print(format_header(["Layer"] + comps, widths))
print("-" * sum(widths))
layer_data: dict[int, dict[str, dict]] = {}
for key, r in results.items():
parts = key.split(".")
if parts[0] != "layers":
continue
li = int(parts[1])
sub = parts[2:]
if sub[0] == "attention":
cname = f"attn.{sub[1]}"
elif sub[0] == "mlp":
cname = f"mlp.{sub[1]}"
else:
continue
layer_data.setdefault(li, {})[cname] = r
for li in sorted(layer_data):
values = [str(li)]
for c in comps:
v = layer_data[li].get(c, {}).get(metric, 0)
values.append(f"{v:.4f}")
print(format_row(values, widths))
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(format_header(["Component", "Mean", "Std", "Min", "Max"], widths))
print("-" * sum(widths))
for name in sorted(groups.keys()):
items = groups[name]
means = [r.get("mean", 0) for r in items]
stds = [r.get("std", 0) for r in items]
mins = [r.get("min", 0) for r in items]
maxs = [r.get("max", 0) for r in items]
g_mean = sum(means) / len(means)
g_std = sum(stds) / len(stds)
g_min = min(mins)
g_max = max(maxs)
print(
format_row(
[
name,
f"{g_mean:.6f}",
f"{g_std:.6f}",
f"{g_min:.6f}",
f"{g_max:.6f}",
],
widths,
)
)
def print_params_summary(results: dict[str, dict]):
total_2d = sum(
r["shape"][0] * r["shape"][1] for r in results.values() if not r.get("is_1d")
)
total_1d = sum(r["shape"][0] for r in results.values() if r.get("is_1d"))
print(f"\n Total 2D params: {total_2d:,}")
print(f" Total 1D params: {total_1d:,}")
print(f" Total params: {total_2d + total_1d:,}")
def main():
parser = argparse.ArgumentParser(
description="SVD effective rank & weight statistics of a model checkpoint."
)
parser.add_argument(
"--ckpt_dir",
type=str,
required=True,
help="Path to checkpoint directory (containing model.safetensors + config.json).",
)
parser.add_argument(
"--compare",
type=str,
nargs="*",
help="Additional checkpoint directories to compare against.",
)
parser.add_argument(
"--no_svd",
action="store_true",
help="Skip SVD analysis, only show weight statistics (mean/std/min/max).",
)
parser.add_argument(
"--output",
type=str,
default=None,
help="Save results as JSON to this path.",
)
args = parser.parse_args()
all_results = {}
def analyze_one(ckpt_dir: str, label: str):
ckpt_dir = Path(ckpt_dir)
weights_path = ckpt_dir / "model.safetensors"
if not weights_path.exists():
print(f"ERROR: {weights_path} not found")
return {}
meta = {}
meta_path = ckpt_dir / "meta.json"
if meta_path.exists():
with open(meta_path) as f:
meta = json.load(f)
print(f"\n{'=' * 70}")
print(f" {label}: {ckpt_dir}")
if meta:
print(
f" Iteration: {meta.get('iteration', '?')}, "
f"Strategy: {meta.get('strategy', '?')}, "
f"nprocs={meta.get('nprocs', '?')}"
)
print(f"{'=' * 70}")
print(f"Loading weights...")
sd = safetensors.torch.load_file(str(weights_path))
print(f" {len(sd)} keys loaded")
weight_keys = [
k
for k in sd
if ".weight" in k and "rotary_embedding" not in k and "freqs_cis" not in k
]
results = {}
if not args.no_svd:
print(f"Computing SVD on {len(weight_keys)} tensors...")
for i, k in enumerate(sorted(weight_keys)):
print(f" [{i + 1}/{len(weight_keys)}] {k:<60s}", end="\r")
results[k] = effective_rank_metrics(sd[k])
print()
else:
print(f"Computing stats on {len(weight_keys)} tensors (no SVD)...")
for i, k in enumerate(sorted(weight_keys)):
t = sd[k]
results[k] = {
"shape": tuple(t.shape),
"is_1d": t.ndim == 1,
"mean": t.float().mean().item(),
"std": t.float().std().item(),
"min": t.float().min().item(),
"max": t.float().max().item(),
}
print_params_summary(results)
if not args.no_svd:
print_component_summary(
results, "\n=== SVD Effective Rank by Component ==="
)
print_layer_grid(results)
print_weight_stats(results)
all_results[label] = results
return results
analyze_one(args.ckpt_dir, "Primary")
if args.compare:
for cdir in args.compare:
analyze_one(cdir, f"Compare_{cdir}")
if args.output:
with open(args.output, "w", encoding="utf-8") as f:
json.dump(all_results, f, indent=2)
print(f"\nResults saved to {args.output}")
if __name__ == "__main__":
main()
+293 -221
View File
@@ -1,36 +1,38 @@
"""HumanEval code generation benchmark. """HumanEval benchmark — functional pipeline design.
Generates n completions per problem, extracts function bodies, executes Pipeline:
against hidden tests, and computes pass@k. load -> generate -> extract -> test -> score -> report
Usage:: Each stage is a pure function (except GPU/CPU-bound I/O stages).
Config is a single dataclass; side effects are isolated at pipeline boundaries.
python scripts/tools/evaluate_humaneval.py --param_path ./params \
--data_path HumanEval.jsonl.gz --output results.json \
--num_samples 200 --temperature 0.8 --max_tokens 512
""" """
import argparse import argparse
import json import json
import os import os
import re import re
import subprocess
import sys
from dataclasses import dataclass
from math import prod from math import prod
from multiprocessing import Process, Queue from typing import Dict, Iterator, List, Optional, Sequence, Tuple
from typing import Dict, List, Optional, Tuple
import numpy as np import numpy as np
import torch import torch
import tqdm import tqdm
from datasets import load_dataset
from astrai.inference import InferenceEngine from astrai.inference import InferenceEngine
from astrai.model import AutoModel from astrai.model import AutoModel
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
HUMANEVAL_URL = ( # ---------------------------------------------------------------------------
"https://github.com/openai/human-eval/raw/master/data/HumanEval.jsonl.gz" # Config
) # ---------------------------------------------------------------------------
_STOP_SEQUENCES = [ HUMANEVAL_HF_DATASET = "openai/openai_humaneval"
STOP_SEQUENCES = [
"\nclass ", "\nclass ",
"\ndef ", "\ndef ",
"\n# ", "\n# ",
@@ -40,43 +42,80 @@ _STOP_SEQUENCES = [
] ]
def _download_humaneval(data_path: str): @dataclass
if os.path.exists(data_path): class EvalConfig:
param_path: str = "./params"
data_path: str = "./humaneval/HumanEval.jsonl"
output: Optional[str] = None
test_only: Optional[str] = None
generate_only: bool = False
num_samples: int = 200
max_tokens: int = 512
temperature: float = 0.8
top_p: float = 0.95
top_k: int = 50
batch_size: int = 32
test_timeout: float = 3.0
test_workers: int = 8
k_values: Tuple[int, ...] = (1, 10, 100)
problem_indices: Optional[List[int]] = None
def download(path: str):
if os.path.exists(path):
return return
import gzip os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
import urllib.request print(f"Downloading HumanEval from HuggingFace ({HUMANEVAL_HF_DATASET}) ...")
ds = load_dataset(HUMANEVAL_HF_DATASET, split="test")
os.makedirs(os.path.dirname(data_path) or ".", exist_ok=True) with open(path, "w", encoding="utf-8") as f:
print(f"Downloading HumanEval from {HUMANEVAL_URL} ...") for item in ds:
tmp = data_path + ".tmp" f.write(json.dumps(item, ensure_ascii=False) + "\n")
urllib.request.urlretrieve(HUMANEVAL_URL, tmp) print(f" saved {len(ds)} problems to {path}")
with gzip.open(tmp, "rb") as f_in:
with open(data_path, "wb") as f_out:
f_out.write(f_in.read())
os.remove(tmp)
print(f" saved to {data_path}")
def _load_problems(data_path: str) -> List[dict]: def load_jsonl(path: str) -> List[dict]:
problems = [] rows = []
with open(data_path, "r", encoding="utf-8") as f: with open(path, encoding="utf-8") as f:
for line in f: for line in f:
line = line.strip() line = line.strip()
if line: if line:
problems.append(json.loads(line)) rows.append(json.loads(line))
return problems return rows
def _extract_function_body(code: str, entry_point: str) -> Optional[str]: def save_json(path: str, data):
"""Extract the function body from a completion.""" with open(path, "w", encoding="utf-8") as f:
json.dump(data, f, indent=2, ensure_ascii=False)
def create_engine(param_path: str, batch_size: int) -> InferenceEngine:
model = AutoModel.from_pretrained(param_path)
tokenizer = AutoTokenizer.from_pretrained(param_path)
model.to(device="cuda", dtype=torch.bfloat16)
return InferenceEngine(
model=model,
tokenizer=tokenizer,
max_batch_size=batch_size,
)
def trim_stop(text: str) -> str:
for stop in STOP_SEQUENCES:
idx = text.find(stop)
if idx != -1:
text = text[:idx]
return text
def extract_body(code: str, entry_point: str) -> Optional[str]:
pattern = rf"def\s+{re.escape(entry_point)}\b[^:]*:" pattern = rf"def\s+{re.escape(entry_point)}\b[^:]*:"
match = re.search(pattern, code) match = re.search(pattern, code)
if not match: if not match:
# Use the full code as-is if we can't find the function
return code return code
body_start = match.end() lines = code[match.end() :].split("\n")
lines = code[body_start:].split("\n")
body_lines = [] body_lines = []
started = False started = False
@@ -94,240 +133,273 @@ def _extract_function_body(code: str, entry_point: str) -> Optional[str]:
body_lines.append(stripped) body_lines.append(stripped)
body = "\n".join(body_lines) body = "\n".join(body_lines)
if not body.strip(): return body if body.strip() else None
return None
return body
def _trim_stop_sequences(text: str) -> str: def deduplicate(seq: Sequence[str]) -> List[str]:
for stop in _STOP_SEQUENCES:
idx = text.find(stop)
if idx != -1:
text = text[:idx]
return text
def _execute_code(problem: dict, completion: str, timeout: float = 3.0) -> bool:
"""Run the completion against hidden tests in a subprocess."""
def _worker(queue, full_code):
try:
namespace = {}
exec(full_code, namespace)
check = namespace.get("check")
if check is None:
queue.put(False)
return
check(namespace.get(problem["entry_point"]))
queue.put(True)
except Exception:
queue.put(False)
full_code = problem["prompt"] + completion + "\n" + problem["test"]
queue: Queue = Queue()
proc = Process(target=_worker, args=(queue, full_code))
proc.start()
proc.join(timeout)
if proc.is_alive():
proc.terminate()
proc.join()
return False
try:
return queue.get_nowait()
except Exception:
return False
def _pass_at_k(n: int, c: int, k: int) -> float:
"""Unbiased estimator of pass@k."""
if n - c < k:
return 1.0
return 1.0 - float(prod(1.0 - k / np.arange(n - c + 1, n + 1)))
def _deduplicate(completions: List[str]) -> List[str]:
seen = set() seen = set()
unique = [] return [x for x in seq if not (x in seen or seen.add(x))]
for c in completions:
if c not in seen:
seen.add(c)
unique.append(c)
return unique
def _generate( def generate_batch(
engine: InferenceEngine, engine: InferenceEngine,
prompt: str, prompt: str,
num_samples: int, n: int,
batch_size: int,
max_tokens: int, max_tokens: int,
temperature: float, temperature: float,
top_p: float, top_p: float,
top_k: int, top_k: int,
batch_size: int,
) -> List[str]: ) -> List[str]:
batches = [prompt] * min(batch_size, num_samples)
completions = [] completions = []
remaining = num_samples remaining = n
while remaining > 0: while remaining > 0:
current = min(batch_size, remaining) current = min(batch_size, remaining)
batch_prompts = batches[:current]
outputs = engine.generate( outputs = engine.generate(
prompt=batch_prompts, prompt=[prompt] * current,
stream=False, stream=False,
max_tokens=max_tokens, max_tokens=max_tokens,
temperature=temperature, temperature=temperature,
top_p=top_p, top_p=top_p,
top_k=top_k, top_k=top_k,
) )
if isinstance(outputs, str): completions.extend(outputs if isinstance(outputs, list) else [outputs])
outputs = [outputs]
completions.extend(outputs)
remaining -= current remaining -= current
return deduplicate(completions)
return _deduplicate(completions)
def evaluate( def extract_completions(
engine: InferenceEngine, raw: Sequence[str],
problems: List[dict], entry_point: str,
num_samples: int, ) -> List[str]:
max_tokens: int, bodies = []
temperature: float, for r in raw:
top_p: float, t = trim_stop(r)
top_k: int, body = extract_body(t, entry_point)
batch_size: int,
k_values: Tuple[int, ...] = (1, 10, 100),
) -> Dict:
results = {}
all_pass_at_k = {k: [] for k in k_values}
for problem in tqdm.tqdm(problems, desc="HumanEval", unit="problem"):
task_id = problem["task_id"]
prompt = problem["prompt"]
entry_point = problem["entry_point"]
raw_completions = _generate(
engine,
prompt,
num_samples,
max_tokens,
temperature,
top_p,
top_k,
batch_size,
)
completions = []
for raw in raw_completions:
trimmed = _trim_stop_sequences(raw)
body = _extract_function_body(trimmed, entry_point)
if body: if body:
completions.append(body) bodies.append(body)
return bodies
passed = 0
for comp in completions:
if _execute_code(problem, comp):
passed += 1
n = len(completions)
c = passed
result = {"task_id": task_id, "n": n, "passed": c}
for k in k_values:
result[f"pass@{k}"] = round(_pass_at_k(n, c, k), 4)
all_pass_at_k[k].append(_pass_at_k(n, c, k))
results[task_id] = result
summary = {}
for k in k_values:
vals = all_pass_at_k[k]
summary[f"pass@{k}"] = round(float(np.mean(vals)), 4)
results["_summary"] = summary
def generate_all(
engine: InferenceEngine,
problems: Sequence[dict],
cfg: EvalConfig,
) -> List[dict]:
results = []
for problem in tqdm.tqdm(problems, desc="Generating", unit="problem"):
raw = generate_batch(
engine,
problem["prompt"],
cfg.num_samples,
cfg.batch_size,
cfg.max_tokens,
cfg.temperature,
cfg.top_p,
cfg.top_k,
)
bodies = extract_completions(raw, problem["entry_point"])
results.append(
dict(
task_id=problem["task_id"],
entry_point=problem["entry_point"],
prompt=problem["prompt"],
test=problem["test"],
completions=bodies,
)
)
return results return results
def main(): def execute_one(args: tuple) -> bool:
parser = argparse.ArgumentParser(description="HumanEval benchmark") full_code, entry_point, timeout = args
parser.add_argument( try:
"--param_path", type=str, default="./params", help="Model directory" r = subprocess.run(
[sys.executable, "-c", full_code],
capture_output=True,
timeout=timeout,
) )
parser.add_argument( return r.returncode == 0
"--data_path", except subprocess.TimeoutExpired:
return False
except Exception:
return False
def test_one(item: dict, cfg: EvalConfig, pool=None) -> Tuple[str, int, int]:
from concurrent.futures import ProcessPoolExecutor
task_id = item["task_id"]
completions = item["completions"]
codes = [
(
item["prompt"] + c + "\n" + item["test"],
item["entry_point"],
cfg.test_timeout,
)
for c in completions
]
n = len(codes)
def _run(p):
return sum(1 for ok in p.map(execute_one, codes) if ok)
if pool is not None:
passed = _run(pool)
else:
with ProcessPoolExecutor(max_workers=cfg.test_workers) as p:
passed = _run(p)
return task_id, n, passed
def test_all(
items: Sequence[dict],
cfg: EvalConfig,
) -> Iterator[Tuple[str, int, int]]:
from concurrent.futures import ProcessPoolExecutor
pool = ProcessPoolExecutor(max_workers=cfg.test_workers)
try:
for item in tqdm.tqdm(items, desc="Testing", unit="problem"):
yield test_one(item, cfg, pool)
finally:
pool.shutdown(wait=True)
def pass_at_k(n: int, c: int, k: int) -> float:
if n - c < k:
return 1.0
return 1.0 - float(prod(1.0 - k / np.arange(n - c + 1, n + 1)))
def score_results(
results: Iterator[Tuple[str, int, int]],
k_values: Tuple[int, ...],
) -> Dict:
"""Score pass@k for each problem.
k values are filtered per-problem: if a problem has n < k samples
(e.g. after deduplication), pass@k is not computed for that problem.
The summary averages only over problems where the k was computed.
"""
scores = {k: [] for k in k_values}
output = {}
for task_id, n, passed in results:
entry = {"task_id": task_id, "n": n, "passed": passed}
for k in k_values:
if k <= n:
pk = round(pass_at_k(n, passed, k), 4)
entry[f"pass@{k}"] = pk
scores[k].append(pk)
else:
entry[f"pass@{k}"] = None
output[task_id] = entry
summary = {}
for k in k_values:
vals = scores[k]
if vals:
summary[f"pass@{k}"] = round(float(np.mean(vals)), 4)
else:
summary[f"pass@{k}"] = None
output["_summary"] = summary
return output
def run_pipeline(cfg: EvalConfig) -> Dict:
if cfg.test_only:
with open(cfg.test_only, encoding="utf-8") as f:
generated = json.load(f)
else:
download(cfg.data_path)
problems = load_jsonl(cfg.data_path)
if cfg.problem_indices:
problems = [problems[i] for i in cfg.problem_indices if i < len(problems)]
engine = create_engine(cfg.param_path, cfg.batch_size)
try:
generated = generate_all(engine, problems, cfg)
finally:
engine.shutdown()
if cfg.output:
mid = cfg.output.replace(".json", "_completions.json")
save_json(mid, generated)
print(f"Completions saved to {mid}")
if cfg.generate_only:
return {}
results = test_all(generated, cfg)
scored = score_results(results, cfg.k_values)
return scored
def parse_args(argv: Optional[List[str]] = None) -> EvalConfig:
p = argparse.ArgumentParser(description="HumanEval benchmark")
p.add_argument("--param_path", type=str, default="./params")
p.add_argument("--data_path", type=str, default="./humaneval/HumanEval.jsonl")
p.add_argument("--output", type=str, default=None)
p.add_argument(
"--test_only",
type=str, type=str,
default="./humaneval/HumanEval.jsonl",
help="HumanEval JSONL file (auto-download if missing)",
)
parser.add_argument("--output", type=str, default=None, help="Output JSON path")
parser.add_argument(
"--num_samples",
type=int,
default=200,
help="Completions per problem",
)
parser.add_argument(
"--max_tokens", type=int, default=512, help="Max generation tokens"
)
parser.add_argument(
"--temperature", type=float, default=0.8, help="Sampling temperature"
)
parser.add_argument("--top_p", type=float, default=0.95, help="Top-p sampling")
parser.add_argument("--top_k", type=int, default=50, help="Top-k sampling")
parser.add_argument(
"--batch_size", type=int, default=1, help="Inference batch size"
)
parser.add_argument(
"--problems",
type=int,
nargs="+",
default=None, default=None,
help="Specific problem indices (0-based)", help="Skip generation, test existing completions JSON",
) )
args = parser.parse_args() p.add_argument(
"--generate_only", action="store_true", help="Only generate, skip testing"
_download_humaneval(args.data_path)
problems = _load_problems(args.data_path)
if args.problems:
problems = [problems[i] for i in args.problems if i < len(problems)]
model = AutoModel.from_pretrained(args.param_path)
tokenizer = AutoTokenizer.from_pretrained(args.param_path)
model.to(device="cuda", dtype=torch.bfloat16)
engine = InferenceEngine(
model=model,
tokenizer=tokenizer,
max_batch_size=args.batch_size,
) )
p.add_argument("--num_samples", type=int, default=200)
p.add_argument("--max_tokens", type=int, default=512)
p.add_argument("--temperature", type=float, default=0.8)
p.add_argument("--top_p", type=float, default=0.95)
p.add_argument("--top_k", type=int, default=50)
p.add_argument("--batch_size", type=int, default=32)
p.add_argument("--test_workers", type=int, default=8)
p.add_argument("--test_timeout", type=float, default=3.0)
p.add_argument("--problems", type=int, nargs="+", default=None)
args = p.parse_args(argv)
results = evaluate( return EvalConfig(
engine=engine, param_path=args.param_path,
problems=problems, data_path=args.data_path,
output=args.output,
test_only=args.test_only,
generate_only=args.generate_only,
num_samples=args.num_samples, num_samples=args.num_samples,
max_tokens=args.max_tokens, max_tokens=args.max_tokens,
temperature=args.temperature, temperature=args.temperature,
top_p=args.top_p, top_p=args.top_p,
top_k=args.top_k, top_k=args.top_k,
batch_size=args.batch_size, batch_size=args.batch_size,
k_values=(1, 10, 100), test_workers=args.test_workers,
test_timeout=args.test_timeout,
problem_indices=args.problems,
) )
summary = results.pop("_summary")
def report(scored: Dict):
summary = scored.pop("_summary", {})
print(f"\n{'=' * 60}") print(f"\n{'=' * 60}")
for k, v in summary.items(): for k, v in summary.items():
if v is not None:
print(f" {k}: {v:.2%}") print(f" {k}: {v:.2%}")
else:
print(f" {k}: N/A")
print(f"{'=' * 60}") print(f"{'=' * 60}")
scored["_summary"] = summary
if args.output:
results["_summary"] = summary
with open(args.output, "w", encoding="utf-8") as f:
json.dump(results, f, indent=2, ensure_ascii=False)
print(f"Results saved to {args.output}")
engine.shutdown() def main():
cfg = parse_args()
scored = run_pipeline(cfg)
report(scored)
if cfg.output:
save_json(cfg.output, scored)
print(f"Results saved to {cfg.output}")
if __name__ == "__main__": if __name__ == "__main__":
+419 -206
View File
@@ -1,248 +1,397 @@
"""IFD (Instruction Following Difficulty) data quality scoring. """IFD (Instruction Following Difficulty) data quality scoring.
Computes IFD scores for instruction-response pairs to guide data selection. IFD = conditional_NLL / unconditional_NLL
IFD = conditional_NLL / unconditional_NLL, where:
- conditional_NLL: average CE loss on response tokens given instruction context - Messages format: plain text concatenation (no chat template)
- unconditional_NLL: average CE loss on response tokens alone - Plain format: raw instr_key + resp_key fields
Higher IFD (close to 1) = instruction provides less help = harder sample. v2 changelog:
Lower IFD (close to 0) = instruction provides strong guidance = easy sample. - Same token set: unconditional pass prefixes resp with a plain-text sentinel
IFD > 1 = instruction misleads the model = likely low-quality data. (default ``\\n``; use ``--sentinel_text ""`` for bos/pad fallback).
Both branches predict the identical N resp tokens.
Usage:: Single-token answers (rl=1) are now supported.
- ctx_len tracked in output
python scripts/eval/ifd.py --param_path ./params \ - skip_reason for None samples (no more silent None)
--input data.jsonl --output data_with_ifd.jsonl \ - --per_token for per-token IFD breakdown
--instr_key instruction --resp_key response
Disable chat template::
python scripts/eval/ifd.py --param_path ./params \
--input data.jsonl --output data_with_ifd.jsonl \
--instr_key instruction --resp_key response \
--no_chat_template
""" """
import argparse import argparse
import glob
import json import json
import os
import statistics
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
import tqdm import tqdm
from astrai.model import AutoModel from astrai.model import AutoModel
from astrai.preprocessing.packing import plan_bfd
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
def compute_ifd( def _pack_bins(pairs, max_len):
model, """BFD bin packing: pack (c+r) into bins of max total length.
tokenizer,
instruction: str, Reuses :func:`plan_bfd` so the BFD heuristic stays single-sourced.
response: str, """
device: str, # Treat each pair as a single sequence of length len(c)+len(r) for
max_len: int = 2048, # planning purposes; plan_bfd works on pure lengths.
use_chat_template: bool = False, fake_sequences = [[0] * (len(c) + len(r)) for c, r in pairs]
) -> dict: plan = plan_bfd(fake_sequences, max_len)
if use_chat_template: return [
return _compute_ifd_with_template( [(i, pairs[i][0], pairs[i][1]) for i in bin_indices] for bin_indices in plan
model, tokenizer, instruction, response, device, max_len ]
def _resolve_sentinel_ids(tokenizer, sentinel_text):
"""Tokenize the sentinel text for the unconditional pass prefix.
Falls back to bos/pad_token_id when sentinel_text is empty or
cannot be encoded.
"""
if sentinel_text:
ids = tokenizer.encode(sentinel_text, add_special_tokens=False)
if ids:
return ids
for attr in ("bos_token_id", "pad_token_id", "eos_token_id"):
tid = getattr(tokenizer, attr, None)
if tid is not None:
return [tid]
return [0]
def _collect_input_files(input_path: str) -> list:
"""Resolve *input_path* to a list of JSONL/JSON files."""
if os.path.isdir(input_path):
files = []
for ext in ("*.jsonl", "*.json"):
files.extend(
sorted(glob.glob(os.path.join(input_path, "**", ext), recursive=True))
) )
return _compute_ifd_raw(model, tokenizer, instruction, response, device, max_len) return files
return sorted(glob.glob(input_path))
def _compute_ifd_raw(model, tokenizer, instruction, response, device, max_len) -> dict: def _load_items(filepath: str) -> list:
instr_ids = tokenizer.encode(instruction) """Load JSONL or JSON (array / single dict) into a list of dicts."""
resp_ids = tokenizer.encode(response) with open(filepath, "r", encoding="utf-8") as f:
if filepath.lower().endswith(".json"):
data = json.load(f)
if isinstance(data, dict):
return [data]
return data
return [json.loads(line) for line in f if line.strip()]
if not resp_ids:
return { @torch.inference_mode()
"L_cond": None, def _score_batch(
"L_uncond": None, pairs, model, device, max_len=2048, sentinel_ids=None, per_token=False
"ifd": None, ):
"error": "empty response", """BFD-packed IFD with text-sentinel-anchored unconditional pass.
Conditional: (ctx + resp[0..i-1]) resp[i], i = 0..N-1
Unconditional: (<sentinel> + resp[0..i-1]) resp[i], i = 0..N-1
Both branches predict the identical N response tokens. A short
plain-text sentinel gives the unconditional pass a prefix so that
every response token can be predicted. Single-token answers (rl=1)
are supported.
"""
if not pairs:
return []
if sentinel_ids is None:
sentinel_ids = [0]
bins = _pack_bins(pairs, max_len)
result = [None] * len(pairs)
# ---- conditional pass (packed, per-document position IDs) ----
for bin_items in bins:
seq_ids = []
global_pos = []
doc_ids = []
doc_offsets = []
for di, (orig_idx, c, r) in enumerate(bin_items):
ctx_len = len(c)
start = len(seq_ids)
item_len = len(c) + len(r)
seq_ids.extend(c)
seq_ids.extend(r)
end = len(seq_ids)
global_pos.extend(range(item_len))
doc_ids.extend([di] * item_len)
doc_offsets.append((start, end, orig_idx, ctx_len))
full_ids = torch.tensor([seq_ids], device=device, dtype=torch.long)
pos_ids = torch.tensor([global_pos], device=device, dtype=torch.long)
seq_len = len(seq_ids)
causal = torch.tril(
torch.ones(seq_len, seq_len, dtype=torch.bool, device=device)
)
doc_t = torch.tensor([doc_ids], device=device)
doc_mask = doc_t.unsqueeze(-1) == doc_t.unsqueeze(-2)
attn_mask = (causal & doc_mask[0]).unsqueeze(0).unsqueeze(0)
logits_full = model(full_ids, position_ids=pos_ids, input_mask=attn_mask)[
"logits"
][0]
for start, end, orig_idx, ctx_len in doc_offsets:
rl = end - start - ctx_len
resp_start = start + ctx_len - 1
resp_logits = logits_full[resp_start : end - 1]
resp_targets = torch.tensor(
seq_ids[start + ctx_len : end], device=device, dtype=torch.long
)
cond_losses = F.cross_entropy(
resp_logits, resp_targets, reduction="none"
).cpu()
result[orig_idx] = {
"_cond_losses": cond_losses,
"_rl": rl,
"_ctx_len": ctx_len,
} }
qa_len = len(instr_ids) + len(resp_ids) # ---- unconditional pass (sentinel-prefixed, batched 2D) ----
if qa_len > max_len: valid_items = [
overflow = qa_len - max_len (
instr_ids = instr_ids[overflow:] i,
result[i]["_rl"],
result[i]["_ctx_len"],
result[i]["_cond_losses"],
pairs[i][1],
)
for i in range(len(pairs))
if result[i] is not None and "_cond_losses" in result[i]
]
if not valid_items:
return result
instr_len = len(instr_ids) valid_items.sort(key=lambda x: -x[1])
resp_len = len(resp_ids) prefix_len = len(sentinel_ids)
max_rl = prefix_len + max(rl for _, rl, _, _, _ in valid_items)
bsz = len(valid_items)
qa_ids = instr_ids + resp_ids u_batch = torch.zeros(bsz, max_rl, dtype=torch.long, device=device)
qa_tensor = torch.tensor([qa_ids], device=device, dtype=torch.long) for ri, (_, rl, _, _, r_ids) in enumerate(valid_items):
u_batch[ri, :prefix_len] = torch.tensor(sentinel_ids, dtype=torch.long)
u_batch[ri, prefix_len : prefix_len + rl] = torch.tensor(
r_ids, dtype=torch.long
)
with torch.inference_mode(): logits_resp = model(u_batch)["logits"]
logits_qa = model(qa_tensor)["logits"][0]
resp_logits = logits_qa[instr_len - 1 : -1] for ri, (orig_idx, rl, ctx_len, cond_losses, _) in enumerate(valid_items):
resp_targets = torch.tensor(resp_ids, device=device, dtype=torch.long) unp_logits = logits_resp[ri, prefix_len - 1 : prefix_len - 1 + rl]
L_cond = F.cross_entropy(resp_logits, resp_targets, reduction="mean").item() unp_targets = u_batch[ri, prefix_len : prefix_len + rl]
uncond_losses = F.cross_entropy(unp_logits, unp_targets, reduction="none").cpu()
resp_tensor = torch.tensor([resp_ids], device=device, dtype=torch.long)
with torch.inference_mode():
logits_resp = model(resp_tensor)["logits"][0]
unp_logits = logits_resp[:-1]
unp_targets = resp_tensor[0, 1:]
L_uncond = F.cross_entropy(unp_logits, unp_targets, reduction="mean").item()
L_cond = cond_losses.mean().item()
L_uncond = uncond_losses.mean().item()
ifd = L_cond / L_uncond if L_uncond > 0 else None ifd = L_cond / L_uncond if L_uncond > 0 else None
return { out = {
"L_cond": round(L_cond, 6), "L_cond": round(L_cond, 6),
"L_uncond": round(L_uncond, 6), "L_uncond": round(L_uncond, 6),
"ifd": round(ifd, 6) if ifd is not None else None, "ifd": round(ifd, 6) if ifd is not None else None,
"instr_len": instr_len, "ctx_len": ctx_len,
"resp_len": resp_len, "resp_len": rl,
"error": None,
} }
if per_token:
per = [
(round(c.item() / u.item(), 6) if u.item() > 0 else None)
for c, u in zip(cond_losses, uncond_losses)
]
out["ifd_per_token"] = per
result[orig_idx] = out
return result
def _compute_ifd_with_template( def _trim(context_ids, resp_ids, max_len):
model, tokenizer, instruction, response, device, max_len """Truncate to fit max_len, keeping response intact if possible."""
) -> dict: if len(resp_ids) > max_len // 2:
instr_prefix = tokenizer.apply_chat_template( resp_ids = resp_ids[: max_len // 2]
[{"role": "user", "content": instruction}], full_ids = context_ids + resp_ids
tokenize=False, if len(full_ids) <= max_len:
add_generation_prompt=True, return context_ids, resp_ids
)
full_text = tokenizer.apply_chat_template(
[
{"role": "user", "content": instruction},
{"role": "assistant", "content": response},
],
tokenize=False,
add_generation_prompt=False,
)
full_ids = tokenizer.encode(full_text)
prefix_ids = tokenizer.encode(instr_prefix)
resp_ids = tokenizer.encode(response)
if not resp_ids:
return {
"L_cond": None,
"L_uncond": None,
"ifd": None,
"error": "empty response",
}
if len(full_ids) > max_len:
overflow = len(full_ids) - max_len overflow = len(full_ids) - max_len
full_ids = full_ids[overflow:] if overflow >= len(context_ids):
prefix_len = len(prefix_ids) - overflow return [], resp_ids[:max_len]
prefix_len = max(0, prefix_len) return context_ids[overflow:], resp_ids
else:
prefix_len = len(prefix_ids)
cond_tensor = torch.tensor([full_ids], device=device, dtype=torch.long)
with torch.inference_mode():
logits_qa = model(cond_tensor)["logits"][0]
resp_start = prefix_len - 1
resp_end = len(full_ids) - 1
if resp_end <= resp_start:
return {
"L_cond": None,
"L_uncond": None,
"ifd": None,
"error": "response truncated entirely",
}
resp_logits = logits_qa[resp_start:resp_end]
resp_targets = torch.tensor(full_ids[prefix_len:], device=device, dtype=torch.long)
L_cond = F.cross_entropy(resp_logits, resp_targets, reduction="mean").item()
resp_tensor = torch.tensor([resp_ids], device=device, dtype=torch.long)
with torch.inference_mode():
logits_resp = model(resp_tensor)["logits"][0]
unp_logits = logits_resp[:-1]
unp_targets = resp_tensor[0, 1:]
L_uncond = F.cross_entropy(unp_logits, unp_targets, reduction="mean").item()
ifd = L_cond / L_uncond if L_uncond > 0 else None
return {
"L_cond": round(L_cond, 6),
"L_uncond": round(L_uncond, 6),
"ifd": round(ifd, 6) if ifd is not None else None,
"instr_len": prefix_len,
"resp_len": len(resp_ids),
"error": None,
}
def process_file( def process_file(
param_path: str,
input_file: str,
output_file: str,
instr_key: str,
resp_key: str,
max_len: int,
use_chat_template: bool = False,
):
device = "cuda" if torch.cuda.is_available() else "cpu"
dtype = torch.bfloat16 if device == "cuda" else torch.float32
model = AutoModel.from_pretrained(param_path)
tokenizer = AutoTokenizer.from_pretrained(param_path)
model.to(device=device, dtype=dtype)
model.eval()
if use_chat_template and tokenizer._chat_template is None:
raise RuntimeError(
"--use_chat_template specified but tokenizer has no chat template. "
"Add a chat_template to tokenizer_config.json or omit the flag."
)
with open(input_file, "r", encoding="utf-8") as f:
data = [json.loads(line) for line in f if line.strip()]
results = []
ifd_values = []
with torch.inference_mode():
for item in tqdm.tqdm(data, desc="Computing IFD", unit="sample"):
instruction = item[instr_key]
response = item[resp_key]
scores = compute_ifd(
model, model,
tokenizer, tokenizer,
instruction, input_file,
response, output_file,
instr_key,
resp_key,
max_len=2048,
data_format="plain",
batch_size=1,
device=None,
sentinel_ids=None,
per_token=False,
max_samples=None,
):
"""Score a single file, write per-sample JSONL, return summary stats."""
if device is None:
device = "cuda" if torch.cuda.is_available() else "cpu"
if sentinel_ids is None:
sentinel_ids = _resolve_sentinel_ids(tokenizer, "\n")
data = _load_items(input_file)
if max_samples and len(data) > max_samples:
import random
data = random.sample(data, max_samples)
results = []
all_ifds = []
buffer = []
label = os.path.splitext(os.path.basename(input_file))[0]
for item in tqdm.tqdm(data, desc=f" {label}", unit="sample", leave=False):
if data_format == "messages":
turns = []
for i, msg in enumerate(item.get("messages", [])):
if msg.get("role") != "assistant":
continue
ctx_text = "\n\n".join(m["content"] for m in item["messages"][:i])
ctx_ids = tokenizer.encode(ctx_text)
resp_ids = tokenizer.encode(msg["content"], add_special_tokens=False)
ctx_ids, resp_ids = _trim(ctx_ids, resp_ids, max_len)
if ctx_ids and resp_ids:
turns.append((ctx_ids, resp_ids))
if not turns:
results.append(
{
**item,
"ifd": None,
"skip_reason": "no valid assistant turns",
"ifd_turns": [],
}
)
continue
buffer.append((item, turns, "messages"))
else:
ctx_ids = tokenizer.encode(item[instr_key], add_special_tokens=False)
resp_ids = tokenizer.encode(item[resp_key], add_special_tokens=False)
ctx_ids, resp_ids = _trim(ctx_ids, resp_ids, max_len)
if not ctx_ids or not resp_ids:
results.append(
{
**item,
"ifd": None,
"ifd_detail": {"skip_reason": "empty ctx or resp"},
}
)
continue
buffer.append((item, [(ctx_ids, resp_ids)], "plain"))
if len(buffer) >= batch_size:
_flush_buffer(
buffer,
results,
all_ifds,
model,
device, device,
max_len, max_len,
use_chat_template=use_chat_template, sentinel_ids,
per_token,
)
if buffer:
_flush_buffer(
buffer, results, all_ifds, model, device, max_len, sentinel_ids, per_token
) )
ifd_values.append(scores["ifd"])
results.append({**item, "ifd": scores["ifd"], "ifd_detail": scores})
with open(output_file, "w", encoding="utf-8") as f: with open(output_file, "w", encoding="utf-8") as f:
for item in results: for item in results:
f.write(json.dumps(item, ensure_ascii=False) + "\n") f.write(json.dumps(item, ensure_ascii=False) + "\n")
valid_ifd = [v for v in ifd_values if v is not None] valid_ifd = [v for v in all_ifds if v is not None]
stats = {
"samples": len(data),
"valid_ifd": len(valid_ifd),
"skipped": len(data) - len(valid_ifd),
}
if valid_ifd: if valid_ifd:
import statistics stats["mean_ifd"] = statistics.mean(valid_ifd)
stats["median_ifd"] = statistics.median(valid_ifd)
if len(valid_ifd) > 1:
stats["stdev_ifd"] = statistics.stdev(valid_ifd)
stats["min_ifd"] = min(valid_ifd)
stats["max_ifd"] = max(valid_ifd)
print(f"\n{'=' * 50}") print(f"\n{'=' * 50}")
print(f" [{label}]")
print(f"{'=' * 50}")
print(f" Samples: {len(data)}") print(f" Samples: {len(data)}")
print(f" Valid IFD: {len(valid_ifd)}") print(f" Valid IFD: {len(valid_ifd)}")
print(f" Skipped: {len(data) - len(valid_ifd)}")
print(f" Mean IFD: {statistics.mean(valid_ifd):.4f}") print(f" Mean IFD: {statistics.mean(valid_ifd):.4f}")
print(f" Median IFD: {statistics.median(valid_ifd):.4f}") print(f" Median IFD: {statistics.median(valid_ifd):.4f}")
if len(valid_ifd) > 1:
print(f" Stdev IFD: {statistics.stdev(valid_ifd):.4f}") print(f" Stdev IFD: {statistics.stdev(valid_ifd):.4f}")
print(f" Min IFD: {min(valid_ifd):.4f}") print(f" Min IFD: {min(valid_ifd):.4f}")
print(f" Max IFD: {max(valid_ifd):.4f}") print(f" Max IFD: {max(valid_ifd):.4f}")
print(f"{'=' * 50}") print(f"{'=' * 50}")
print(f" Results saved to {output_file}")
return stats
print(f"Results saved to {output_file}")
def _flush_buffer(
buffer, results, all_ifds, model, device, max_len, sentinel_ids, per_token
):
all_pairs = []
indices = []
for item, turns, fmt in buffer:
start = len(all_pairs)
all_pairs.extend(turns)
indices.append((item, turns, fmt, start, len(all_pairs)))
raw = _score_batch(
all_pairs,
model,
device,
max_len,
sentinel_ids=sentinel_ids,
per_token=per_token,
)
for item, turns, fmt, start, end in indices:
turn_scores = raw[start:end]
if fmt == "messages":
valid = [
s for s in turn_scores if s is not None and s.get("ifd") is not None
]
if not valid:
results.append({**item, "ifd": None, "ifd_turns": turn_scores})
else:
avg = sum(s["ifd"] for s in valid) / len(valid)
all_ifds.append(avg)
results.append(
{
**item,
"ifd": avg,
"ifd_detail": valid[0] if len(valid) == 1 else None,
"ifd_turns": turn_scores,
}
)
else:
score = turn_scores[0]
all_ifds.append(score.get("ifd"))
results.append({**item, "ifd": score.get("ifd"), "ifd_detail": score})
buffer.clear()
def main(): def main():
@@ -250,43 +399,107 @@ def main():
description="Compute IFD scores for instruction-response data" description="Compute IFD scores for instruction-response data"
) )
parser.add_argument("--param_path", type=str, required=True, help="Model directory") parser.add_argument("--param_path", type=str, required=True, help="Model directory")
parser.add_argument("--input", type=str, required=True, help="Input JSONL file")
parser.add_argument("--output", type=str, required=True, help="Output JSONL file")
parser.add_argument( parser.add_argument(
"--instr_key", "--input_path",
type=str, type=str,
default="instruction", required=True,
help="Key for instruction field", help="Input file, glob pattern, or directory.",
) )
parser.add_argument( parser.add_argument(
"--resp_key", "--output_dir",
type=str, type=str,
default="response", required=True,
help="Key for response field", help="Directory for output files (summary.json + per-file JSONL).",
)
parser.add_argument("--max_len", type=int, default=2048, help="Max token length")
parser.add_argument(
"--format",
type=str,
default="plain",
choices=["plain", "messages"],
help="Input format",
) )
parser.add_argument( parser.add_argument(
"--max_len", "--instr_key", type=str, default="instruction", help="Key for instruction field"
type=int,
default=2048,
help="Max token length (instruction truncated to fit)",
) )
parser.add_argument( parser.add_argument(
"--no_chat_template", "--resp_key", type=str, default="response", help="Key for response field"
)
parser.add_argument(
"--batch_size", type=int, default=8, help="Batch size for model forward passes"
)
parser.add_argument("--device", type=str, default=None, help="Device (e.g. cuda:0)")
parser.add_argument(
"--dtype",
type=str,
default="bfloat16" if torch.cuda.is_available() else "float32",
help="Torch dtype",
)
parser.add_argument(
"--sentinel_text",
type=str,
default="\n",
help='Plain-text prefix for unconditional pass (default: "\\n"). Use "" for bos/pad fallback.',
)
parser.add_argument(
"--per_token",
action="store_true", action="store_true",
default=False, help="Include per-token IFD breakdown in output",
help="Disable chat template, use raw text concatenation", )
parser.add_argument(
"--max_samples",
type=int,
default=None,
help="Maximum number of samples per file (random subsample). Default: all.",
) )
args = parser.parse_args() args = parser.parse_args()
process_file( if args.device is None:
args.param_path, args.device = "cuda" if torch.cuda.is_available() else "cpu"
args.input, dtype = getattr(torch, args.dtype)
args.output,
args.instr_key, print(f"Loading model from {args.param_path} ...")
args.resp_key, model = AutoModel.from_pretrained(args.param_path)
args.max_len, tokenizer = AutoTokenizer.from_pretrained(args.param_path)
use_chat_template=not args.no_chat_template, model.to(device=args.device, dtype=dtype)
model.eval()
sentinel_ids = _resolve_sentinel_ids(tokenizer, args.sentinel_text)
input_files = _collect_input_files(args.input_path)
if not input_files:
print(f"No input files found at {args.input_path}")
return
print(f"Found {len(input_files)} file(s) to evaluate")
os.makedirs(args.output_dir, exist_ok=True)
all_stats = {}
for filepath in input_files:
label = os.path.splitext(os.path.basename(filepath))[0]
output_file = os.path.join(args.output_dir, f"{label}_ifd.jsonl")
stats = process_file(
model=model,
tokenizer=tokenizer,
input_file=filepath,
output_file=output_file,
instr_key=args.instr_key,
resp_key=args.resp_key,
max_len=args.max_len,
data_format=args.format,
batch_size=args.batch_size,
device=args.device,
sentinel_ids=sentinel_ids,
per_token=args.per_token,
max_samples=args.max_samples,
) )
all_stats[label] = stats
summary_path = os.path.join(args.output_dir, "summary.json")
with open(summary_path, "w", encoding="utf-8") as f:
json.dump(all_stats, f, ensure_ascii=False, indent=2)
print(f"\nSummary saved to {summary_path}")
if __name__ == "__main__": if __name__ == "__main__":
+10 -17
View File
@@ -5,7 +5,7 @@ Supports all IFEval constraint types except language detection.
Usage:: Usage::
python scripts/tools/evaluate_ifeval.py --param_path ./params \ python scripts/eval/evaluate_ifeval.py --param_path ./params \
--data_path ifeval.jsonl --output results.json \ --data_path ifeval.jsonl --output results.json \
--temperature 0.1 --max_tokens 512 --temperature 0.1 --max_tokens 512
""" """
@@ -14,21 +14,17 @@ import argparse
import json import json
import os import os
import re import re
import urllib.request
from typing import Callable, Dict, List, Optional from typing import Callable, Dict, List, Optional
import torch import torch
import tqdm import tqdm
from datasets import load_dataset
from astrai.inference import InferenceEngine from astrai.inference import InferenceEngine
from astrai.model import AutoModel from astrai.model import AutoModel
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
IFEVAL_URL = ( IFEVAL_HF_DATASET = "google/IFEval"
"https://raw.githubusercontent.com/google-research/"
"google-research/master/instruction_following_eval/data/input_data.jsonl"
)
CONSTRAINT_VERIFIERS: Dict[str, Callable[[str, dict], bool]] = {} CONSTRAINT_VERIFIERS: Dict[str, Callable[[str, dict], bool]] = {}
@@ -310,15 +306,12 @@ def download_ifeval(data_path: str):
if os.path.exists(data_path): if os.path.exists(data_path):
return return
os.makedirs(os.path.dirname(data_path) or ".", exist_ok=True) os.makedirs(os.path.dirname(data_path) or ".", exist_ok=True)
print(f"Downloading IFEval from {IFEVAL_URL} ...") print(f"Downloading IFEval from HuggingFace ({IFEVAL_HF_DATASET}) ...")
tmp = data_path + ".tmp" ds = load_dataset(IFEVAL_HF_DATASET, split="train")
urllib.request.urlretrieve(IFEVAL_URL, tmp) with open(data_path, "w", encoding="utf-8") as f:
with open(tmp, "rb") as f_in: for item in ds:
content = f_in.read() f.write(json.dumps(item, ensure_ascii=False) + "\n")
with open(data_path, "wb") as f_out: print(f" saved {len(ds)} items to {data_path}")
f_out.write(content)
os.remove(tmp)
print(f" saved to {data_path}")
def load_problems(data_path: str) -> List[dict]: def load_problems(data_path: str) -> List[dict]:
@@ -571,7 +564,7 @@ def main():
print(f" Unsupported: {summary['unsupported_constraints']}") print(f" Unsupported: {summary['unsupported_constraints']}")
print(f"{'=' * 60}") print(f"{'=' * 60}")
print(f"\nPer-type accuracy:") print("\nPer-type accuracy:")
for inst_id, stats in sorted(summary["per_type_accuracy"].items()): for inst_id, stats in sorted(summary["per_type_accuracy"].items()):
print( print(
f" {inst_id:50s} {stats['accuracy']:.2%} " f" {inst_id:50s} {stats['accuracy']:.2%} "
+91 -55
View File
@@ -4,18 +4,18 @@ import argparse
import csv import csv
import json import json
import os import os
import shutil import random
import tarfile from collections import defaultdict
import requests
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
import tqdm import tqdm
from datasets import load_dataset
from astrai.model import AutoModel from astrai.model import AutoModel
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
MMLU_URL = "https://people.eecs.berkeley.edu/~hendrycks/data.tar" MMLU_HF_DATASET = "cais/mmlu"
MMLU_SUBJECTS = [ MMLU_SUBJECTS = [
"abstract_algebra", "abstract_algebra",
"anatomy", "anatomy",
@@ -77,38 +77,40 @@ MMLU_SUBJECTS = [
] ]
def _download_and_extract(url: str, data_dir: str): def _write_subject_csv(data_dir: str, split: str, subject: str, rows: list[dict]):
tar_path = os.path.join(data_dir, "data.tar") split_dir = os.path.join(data_dir, split)
os.makedirs(data_dir, exist_ok=True) os.makedirs(split_dir, exist_ok=True)
print(f"Downloading MMLU data from {url}...") path = os.path.join(split_dir, f"{subject}_{split}.csv")
resp = requests.get(url, stream=True, timeout=300) with open(path, "w", encoding="utf-8", newline="") as f:
resp.raise_for_status() writer = csv.writer(f)
total = int(resp.headers.get("content-length", 0)) for row in rows:
with tqdm.tqdm(total=total, unit="B", unit_scale=True, desc=" Download") as bar: writer.writerow(row)
with open(tar_path, "wb") as f:
for chunk in resp.iter_content(chunk_size=8192):
f.write(chunk)
bar.update(len(chunk))
print("Extracting...")
with tarfile.open(tar_path, "r") as tf:
tf.extractall(data_dir)
os.remove(tar_path)
def download_mmlu(data_dir: str): def download_mmlu(data_dir: str):
_download_and_extract(MMLU_URL, data_dir) print(f"Downloading MMLU from HuggingFace ({MMLU_HF_DATASET}) ...")
src = os.path.join(data_dir, "data") letters = ("A", "B", "C", "D")
if os.path.exists(src): split_map = {"dev": "dev", "val": "validation", "test": "test"}
for item in os.listdir(src): for local_split, hf_split in split_map.items():
src_item = os.path.join(src, item) ds = load_dataset(MMLU_HF_DATASET, "all", split=hf_split)
dst_item = os.path.join(data_dir, item) grouped: dict[str, list[dict]] = defaultdict(list)
if os.path.exists(dst_item): for item in tqdm.tqdm(ds, desc=f" {local_split}", leave=False):
if os.path.isdir(dst_item): subject = item["subject"]
shutil.rmtree(dst_item) choices = item["choices"]
else: ans_letter = letters[item["answer"]]
os.remove(dst_item) grouped[subject].append(
os.rename(src_item, dst_item) [
os.rmdir(src) item["question"],
f"A){choices[0]}",
f"B){choices[1]}",
f"C){choices[2]}",
f"D){choices[3]}",
ans_letter,
]
)
for subject, rows in grouped.items():
_write_subject_csv(data_dir, local_split, subject, rows)
print(f" {local_split}: {len(ds)} items, {len(grouped)} subjects")
print(f"MMLU data saved to {data_dir}") print(f"MMLU data saved to {data_dir}")
@@ -139,17 +141,12 @@ def load_csv(path: str) -> list[dict]:
return data return data
def build_prompt( def build_prompt(question: str, choices: dict, subject: str) -> str:
question: str, choices: dict, subject: str, n_shot: int, dev_data: list[dict] """Build the raw question prompt (without few-shot examples).
) -> str:
prompt = "" Few-shot examples are handled by ``apply_chat`` to avoid duplication.
if n_shot > 0 and dev_data: """
prompt = f"The following are multiple choice questions (with answers) about {subject}.\n\n" prompt = f"The following are multiple choice questions (with answers) about {subject}.\n\n"
for item in dev_data[:n_shot]:
prompt += f"Question: {item['question']}\n"
for k in ("A", "B", "C", "D"):
prompt += f"{k}. {item[k]}\n"
prompt += f"Answer: {item['answer']}\n\n"
prompt += f"Question: {question}\n" prompt += f"Question: {question}\n"
for k in ("A", "B", "C", "D"): for k in ("A", "B", "C", "D"):
prompt += f"{k}. {choices[k]}\n" prompt += f"{k}. {choices[k]}\n"
@@ -158,19 +155,22 @@ def build_prompt(
def apply_chat( def apply_chat(
tokenizer, raw_prompt: str, n_shot: int, dev_data: list[dict] | None tokenizer,
raw_prompt: str,
n_shot: int,
dev_data: list[dict] | None,
subject: str = "",
) -> str: ) -> str:
"""Wrap raw MMLU prompt in the model's chat template format. """Wrap raw MMLU prompt in the model's chat template format.
For few-shot, prepend example Q&A pairs as a second user/assistant exchange. For few-shot, prepend example Q&A pairs as user/assistant exchanges.
Few-shot examples use the same subject preamble as the test question to
keep the format consistent.
""" """
messages = [] messages = []
if n_shot > 0 and dev_data: if n_shot > 0 and dev_data:
for item in dev_data[:n_shot]: for item in dev_data[:n_shot]:
q = f"Question: {item['question']}\n" q = build_prompt(item["question"], item, subject)
for k in ("A", "B", "C", "D"):
q += f"{k}. {item[k]}\n"
q += "Answer:"
messages.append({"role": "user", "content": q}) messages.append({"role": "user", "content": q})
messages.append({"role": "assistant", "content": item["answer"]}) messages.append({"role": "assistant", "content": item["answer"]})
messages.append({"role": "user", "content": raw_prompt}) messages.append({"role": "user", "content": raw_prompt})
@@ -206,6 +206,25 @@ def choice_logprob(
return score return score
def _permute_choices(item: dict, rng: random.Random) -> tuple[dict, str]:
"""Shuffle the option order of a question.
Returns ``(permuted_item, new_answer_letter)``. The question text and
the *content* of each choice are unchanged; only which letter (A/B/C/D)
maps to which content is shuffled. This neutralises the model's
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"]}
for new_letter, orig_letter in zip(letters, perm):
permuted[new_letter] = item[orig_letter]
new_answer = letters[perm.index(item["answer"])]
return permuted, new_answer
def evaluate_subject( def evaluate_subject(
model, model,
tokenizer, tokenizer,
@@ -214,20 +233,24 @@ def evaluate_subject(
dev_data: list[dict] | None, dev_data: list[dict] | None,
device: str, device: str,
n_shot: int, n_shot: int,
seed: int = 0,
) -> tuple[float, int, int]: ) -> tuple[float, int, int]:
rng = random.Random(seed) if seed >= 0 else None
correct = 0 correct = 0
total = 0 total = 0
for item in tqdm.tqdm(test_data, desc=f"{subject:40s}", leave=False): for item in tqdm.tqdm(test_data, desc=f"{subject:40s}", leave=False):
raw_prompt = build_prompt( if rng is not None:
item["question"], item, subject, n_shot, dev_data or [] permuted, answer = _permute_choices(item, rng)
) else:
context = apply_chat(tokenizer, raw_prompt, n_shot, dev_data or []) permuted, answer = item, item["answer"]
raw_prompt = build_prompt(permuted["question"], permuted, subject)
context = apply_chat(tokenizer, raw_prompt, n_shot, dev_data or [], subject)
context_ids = tokenizer.encode(context) context_ids = tokenizer.encode(context)
scores = { scores = {
c: choice_logprob(model, tokenizer, context_ids, c, device) c: choice_logprob(model, tokenizer, context_ids, c, device)
for c in ("A", "B", "C", "D") for c in ("A", "B", "C", "D")
} }
if max(scores, key=scores.get) == item["answer"]: if max(scores, key=scores.get) == answer:
correct += 1 correct += 1
total += 1 total += 1
return correct / total, correct, total return correct / total, correct, total
@@ -262,6 +285,12 @@ def main():
default="bfloat16" if torch.cuda.is_available() else "float32", default="bfloat16" if torch.cuda.is_available() else "float32",
help="Torch dtype", help="Torch dtype",
) )
parser.add_argument(
"--seed",
type=int,
default=0,
help="Seed for option permutation (0 to enable, -1 to disable)",
)
args = parser.parse_args() args = parser.parse_args()
if args.download or not os.path.exists(args.data_dir): if args.download or not os.path.exists(args.data_dir):
@@ -293,7 +322,14 @@ def main():
test_data = load_csv(test_path) test_data = load_csv(test_path)
acc, corr, tot = evaluate_subject( acc, corr, tot = evaluate_subject(
model, tokenizer, subject, test_data, dev_data, device, args.n_shot model,
tokenizer,
subject,
test_data,
dev_data,
device,
args.n_shot,
seed=args.seed,
) )
results[subject] = {"accuracy": round(acc, 4), "correct": corr, "total": tot} results[subject] = {"accuracy": round(acc, 4), "correct": corr, "total": tot}
total_correct += corr total_correct += corr
+415 -62
View File
@@ -1,5 +1,9 @@
import argparse import argparse
import glob
import json import json
import os
import statistics
from typing import Dict, List, Optional, Tuple
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
@@ -9,95 +13,400 @@ from astrai.model import AutoModel
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
def process_file( def _collect_input_files(input_path: str) -> List[str]:
param_path: str, input_file: str, output_file: str, batch_size: int, text_key: str """Resolve *input_path* to a list of JSONL/JSON files."""
): if os.path.isdir(input_path):
# Load model and tokenizer files = []
model = AutoModel.from_pretrained(param_path) for ext in ("*.jsonl", "*.json"):
tokenizer = AutoTokenizer.from_pretrained(param_path) files.extend(
model.to(device="cuda", dtype=torch.bfloat16) sorted(glob.glob(os.path.join(input_path, "**", ext), recursive=True))
)
return files
return sorted(glob.glob(input_path))
with open(input_file, "r", encoding="utf-8") as f:
input_data = [json.loads(line) for line in f]
texts = [item[text_key] for item in input_data] def _load_items(filepath: str) -> List[dict]:
"""Load JSONL or JSON (array / single dict) into a list of dicts."""
with open(filepath, "r", encoding="utf-8") as f:
if filepath.lower().endswith(".json"):
data = json.load(f)
if isinstance(data, dict):
return [data]
return data
return [json.loads(line) for line in f if line.strip()]
# Encode all texts
print(f"Encoding {len(texts)} texts...")
encoded_texts = [tokenizer.encode(text) for text in texts]
output_data = [] def _encode_batch(
total_batches = (len(encoded_texts) + batch_size - 1) // batch_size tokenizer: AutoTokenizer, texts: List[str], max_length: int
) -> Tuple[List[List[int]], List[List[int]]]:
"""Encode *texts* and return (token_ids, attention_masks).
for i in tqdm.tqdm( Each sequence is left-aligned and padded to the batch max length.
range(0, len(encoded_texts), batch_size), """
total=total_batches, encoded = [tokenizer.encode(t)[:max_length] for t in texts]
desc="Computing perplexity", if not encoded:
): return [], []
batch_encoded = encoded_texts[i : i + batch_size] max_len = max(len(seq) for seq in encoded)
batch_texts = texts[i : i + batch_size]
# Find max length in batch and pad
max_len = max(len(seq) for seq in batch_encoded)
padded_ids = [] padded_ids = []
masks = [] masks = []
for seq in encoded:
for seq in batch_encoded:
pad_len = max_len - len(seq) pad_len = max_len - len(seq)
padded_seq = seq + [tokenizer.pad_id] * pad_len padded_ids.append(seq + [tokenizer.pad_id] * pad_len)
mask = [True] * len(seq) + [False] * pad_len masks.append([1] * len(seq) + [0] * pad_len)
padded_ids.append(padded_seq) return padded_ids, masks
masks.append(mask)
# Convert to tensors
input_ids = torch.tensor(padded_ids, device="cuda", dtype=torch.long)
input_mask = torch.tensor(masks, device="cuda", dtype=torch.bool)
# Compute perplexity def _compute_batch(
output = model(input_ids, input_mask=input_mask) model,
logits = output["logits"] input_ids: torch.Tensor,
attention_mask: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Forward pass and return (log_probs, valid_mask) of shape [B, S-1].
# Shift for causal language modeling log_probs[i, j] = log P(token j+1 | tokens 0..j)
shifted_logits = logits[:, :-1, :] # [batch_size, seq_len-1, vocab_size] """
shifted_input_ids = input_ids[:, 1:] # [batch_size, seq_len-1] output = model(input_ids, input_mask=attention_mask)
shifted_mask = input_mask[:, 1:] # [batch_size, seq_len-1] logits = output["logits"][:, :-1, :] # [B, S-1, V]
targets = input_ids[:, 1:] # [B, S-1]
valid = attention_mask[:, 1:].float() # [B, S-1]
# Compute cross entropy loss log_probs = F.log_softmax(logits.float(), dim=-1) # [B, S-1, V]
loss = F.cross_entropy( token_log_probs = log_probs.gather(2, targets.unsqueeze(-1)).squeeze(-1) # [B, S-1]
shifted_logits.flatten(0, 1),
shifted_input_ids.flatten(0, 1), return token_log_probs, valid
reduction="none",
def _token_type(token_id: int, stop_ids: frozenset, decode_fn) -> str:
"""Classify a token into a coarse type for analysis.
*stop_ids* is a pre-built set of special token IDs.
*decode_fn* is ``tokenizer.decode`` (or a wrapper) for single-token
decoding.
"""
if token_id in stop_ids:
return "special"
decoded = decode_fn([token_id], skip_special_tokens=True)
if any("\u4e00" <= ch <= "\u9fff" for ch in decoded):
return "cjk"
if any(ord(ch) > 127 for ch in decoded):
return "non_ascii"
return "ascii"
def _percentiles(values: List[float]) -> Dict[str, float]:
"""Compute common percentiles from a list of floats.
Uses linear interpolation between closest ranks (same convention
as NumPy's default).
"""
if not values:
return {}
sorted_vals = sorted(values)
n = len(sorted_vals)
def _pct(p: float) -> float:
if n == 1:
return sorted_vals[0]
k = p * (n - 1)
f = int(k)
c = min(f + 1, n - 1)
return sorted_vals[f] + (sorted_vals[c] - sorted_vals[f]) * (k - f)
return {
"p50": _pct(0.50),
"p90": _pct(0.90),
"p95": _pct(0.95),
"p99": _pct(0.99),
}
class LossAccumulator:
"""Accumulate per-token losses with optional streaming mode.
When *stream* is True (token_level=False), losses are not kept
in memory individually only a running sum/count and a histogram
(for approximate percentiles) are maintained. When *stream* is
False, all losses are retained for exact statistics and per-record
output.
"""
_HIST_BINS = 1000
_HIST_MAX = 20.0 # clamp losses above this for histogram
def __init__(self, stream: bool):
self.stream = stream
self.losses: List[float] = [] if not stream else []
self.total: float = 0.0
self.count: int = 0
self.hist = torch.zeros(self._HIST_BINS, dtype=torch.long)
# per-type losses (only populated when not streaming)
self.by_type: Dict[str, List[float]] = {}
self.type_total: Dict[str, float] = {}
self.type_count: Dict[str, int] = {}
def add(self, losses: List[float]):
self.total += sum(losses)
self.count += len(losses)
if self.stream:
clamped = [min(max(l, 0.0), self._HIST_MAX) for l in losses]
idx = torch.tensor(clamped) / self._HIST_MAX * (self._HIST_BINS - 1)
self.hist += torch.bincount(
idx.long().clamp(0, self._HIST_BINS - 1),
minlength=self._HIST_BINS,
) )
else:
self.losses.extend(losses)
loss = loss.view(shifted_input_ids.shape) # [batch_size, seq_len-1] def add_typed(self, ttype: str, losses: List[float]):
loss = loss * shifted_mask if not self.stream:
sentence_loss = loss.sum(dim=1) / shifted_mask.sum(dim=1).clamp(min=1) self.by_type.setdefault(ttype, []).extend(losses)
perplexity = torch.exp(sentence_loss) # [batch_size] self.type_total[ttype] = self.type_total.get(ttype, 0.0) + sum(losses)
self.type_count[ttype] = self.type_count.get(ttype, 0) + len(losses)
for text, ppl in zip(batch_texts, perplexity): def stats(self) -> Dict:
output_data.append({text_key: text, "ppl": float(ppl.item())}) result: Dict = {}
if self.count == 0:
return result
mean_loss = self.total / self.count
result["overall"] = {
"num_tokens": self.count,
"mean_loss": mean_loss,
"ppl": float(torch.exp(torch.tensor(mean_loss))),
}
if self.stream:
result["overall"].update(self._hist_percentiles())
else:
result["overall"]["median_loss"] = statistics.median(self.losses)
result["overall"].update(_percentiles(self.losses))
# Write results if self.type_count:
result["by_token_type"] = {}
for ttype in sorted(self.type_count.keys()):
cnt = self.type_count[ttype]
tmean = self.type_total[ttype] / cnt
entry: Dict = {
"num_tokens": cnt,
"mean_loss": tmean,
"ppl": float(torch.exp(torch.tensor(tmean))),
}
if not self.stream and ttype in self.by_type:
entry["median_loss"] = statistics.median(self.by_type[ttype])
entry.update(_percentiles(self.by_type[ttype]))
result["by_token_type"][ttype] = entry
return result
def _hist_percentiles(self) -> Dict[str, float]:
"""Approximate percentiles from the histogram."""
total = self.hist.sum().item()
if total == 0:
return {}
cum = torch.cumsum(self.hist.float(), dim=0)
result = {}
for label, p in [("p50", 0.5), ("p90", 0.9), ("p95", 0.95), ("p99", 0.99)]:
target = p * total
idx = int(torch.searchsorted(cum, target).item())
idx = min(idx, self._HIST_BINS - 1)
result[label] = (idx + 0.5) / self._HIST_BINS * self._HIST_MAX
return result
def process_file(
model,
tokenizer: AutoTokenizer,
items: List[dict],
text_key: str,
batch_size: int,
max_length: int,
token_level: bool,
max_samples: Optional[int],
output_file: Optional[str],
label: str,
device: str = "cuda",
) -> Dict:
"""Evaluate a single dataset (list of items), return summary stats.
If *token_level* is True and *output_file* is set, per-record token_ids
and log_probs are written as JSONL alongside the summary.
"""
if max_samples and len(items) > max_samples:
import random
items = random.sample(items, max_samples)
texts = [item[text_key] for item in items if text_key in item]
print(f" [{label}] {len(texts)} samples, text_key='{text_key}'")
acc = LossAccumulator(stream=not token_level)
per_sample: List[dict] = []
if token_level:
stop_ids = frozenset(tokenizer.stop_ids)
decode_fn = tokenizer.decode
num_batches = (len(texts) + batch_size - 1) // batch_size
for i in tqdm.tqdm(
range(0, len(texts), batch_size),
total=num_batches,
desc=f" {label}",
leave=False,
):
batch_texts = texts[i : i + batch_size]
padded_ids, masks = _encode_batch(tokenizer, batch_texts, max_length)
input_ids = torch.tensor(padded_ids, device=device, dtype=torch.long)
attention_mask = torch.tensor(masks, device=device, dtype=torch.bool)
token_log_probs, valid = _compute_batch(model, input_ids, attention_mask)
for b in range(len(batch_texts)):
seq_len = int(valid[b].sum().item())
lps = token_log_probs[b, :seq_len].tolist()
losses = [-lp for lp in lps]
acc.add(losses)
if token_level:
# log_probs correspond to positions 1..seq_len (predicted
# from position 0..seq_len-1), so token_ids must skip BOS
# at position 0 to stay aligned with log_probs.
ids = padded_ids[b][1 : seq_len + 1]
per_sample.append(
{
"text": batch_texts[b][:200],
"token_ids": ids,
"log_probs": [round(lp, 4) for lp in lps],
"ppl": float(torch.exp(torch.tensor(statistics.mean(losses))))
if losses
else None,
}
)
typed_losses: Dict[str, List[float]] = {}
for tid, loss in zip(ids, losses):
ttype = _token_type(tid, stop_ids, decode_fn)
typed_losses.setdefault(ttype, []).append(loss)
for ttype, tl in typed_losses.items():
acc.add_typed(ttype, tl)
stats = acc.stats()
if token_level and output_file:
with open(output_file, "w", encoding="utf-8") as f: with open(output_file, "w", encoding="utf-8") as f:
for item in output_data: for item in per_sample:
f.write(json.dumps(item, ensure_ascii=False) + "\n") f.write(json.dumps(item, ensure_ascii=False) + "\n")
print(f"Perplexity computation complete. Results saved to {output_file}") return stats
def print_stats(label: str, stats: Dict):
"""Pretty-print summary statistics."""
print(f"\n{'=' * 60}")
print(f" {label}")
print(f"{'=' * 60}")
ov = stats.get("overall", {})
if ov:
print(f" tokens: {ov['num_tokens']:,}")
print(f" mean loss: {ov['mean_loss']:.4f}")
if "median_loss" in ov:
print(f" median loss: {ov['median_loss']:.4f}")
print(f" ppl: {ov['ppl']:.2f}")
if "p50" in ov:
print(
f" p50/p90/p95/p99: "
f"{ov['p50']:.2f} / {ov['p90']:.2f} / {ov['p95']:.2f} / {ov['p99']:.2f}"
)
by_type = stats.get("by_token_type", {})
if by_type:
print(f"\n by token type:")
print(f" {'type':<12} {'count':>8} {'mean_loss':>10} {'ppl':>8}")
print(f" {'-' * 12} {'-' * 8} {'-' * 10} {'-' * 8}")
for ttype, s in by_type.items():
print(
f" {ttype:<12} {s['num_tokens']:>8,} "
f"{s['mean_loss']:>10.4f} {s['ppl']:>8.2f}"
)
def main(
param_path: str,
input_path: str,
output_dir: str,
text_key: str,
batch_size: int,
max_length: int,
token_level: bool,
max_samples: Optional[int],
device: str = "cuda",
dtype: str = "bfloat16",
):
print(f"Loading model from {param_path} ...")
model = AutoModel.from_pretrained(param_path)
tokenizer = AutoTokenizer.from_pretrained(param_path)
torch_dtype = getattr(torch, dtype)
model.to(device=device, dtype=torch_dtype)
model.eval()
input_files = _collect_input_files(input_path)
if not input_files:
print(f"No input files found at {input_path}")
return
print(f"Found {len(input_files)} file(s) to evaluate")
os.makedirs(output_dir, exist_ok=True)
all_stats = {}
for filepath in input_files:
label = os.path.splitext(os.path.basename(filepath))[0]
items = _load_items(filepath)
if not items:
print(f" [{label}] empty, skipping")
continue
token_output = (
os.path.join(output_dir, f"{label}_tokens.jsonl") if token_level else None
)
stats = process_file(
model=model,
tokenizer=tokenizer,
items=items,
text_key=text_key,
batch_size=batch_size,
max_length=max_length,
token_level=token_level,
max_samples=max_samples,
output_file=token_output,
label=label,
device=device,
)
all_stats[label] = stats
print_stats(label, stats)
if token_output:
print(f" token-level output: {token_output}")
summary_path = os.path.join(output_dir, "summary.json")
with open(summary_path, "w", encoding="utf-8") as f:
json.dump(all_stats, f, ensure_ascii=False, indent=2)
print(f"\nSummary saved to {summary_path}")
if __name__ == "__main__": if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Run perplexity with a Khaosz model.") parser = argparse.ArgumentParser(
description="Perplexity and token-level loss evaluation on JSONL/JSON data."
)
parser.add_argument( parser.add_argument(
"--param_path", type=str, required=True, help="Path to the model directory." "--param_path", type=str, required=True, help="Path to the model directory."
) )
parser.add_argument( parser.add_argument(
"--input_file", type=str, required=True, help="Path to the input file." "--input_path",
type=str,
required=True,
help="Path to input file, glob pattern, or directory.",
) )
parser.add_argument( parser.add_argument(
"--output_file", type=str, required=True, help="Path to the output file." "--output_dir",
) type=str,
parser.add_argument( required=True,
"--batch_size", type=int, default=4, help="Batch size for evaluation." help="Directory for output files (summary.json + per-file token JSONL).",
) )
parser.add_argument( parser.add_argument(
"--text_key", "--text_key",
@@ -105,7 +414,51 @@ if __name__ == "__main__":
default="text", default="text",
help="Key for the text field in the input data.", help="Key for the text field in the input data.",
) )
parser.add_argument(
"--batch_size", type=int, default=4, help="Batch size for evaluation."
)
parser.add_argument(
"--max_length",
type=int,
default=2048,
help="Maximum sequence length (tokens). Longer sequences are truncated.",
)
parser.add_argument(
"--token_level",
action="store_true",
help="Store per-token log_probs and token type analysis. "
"Default: off (only aggregate stats).",
)
parser.add_argument(
"--max_samples",
type=int,
default=None,
help="Maximum number of samples per file (random subsample). Default: all.",
)
parser.add_argument(
"--device",
type=str,
default="cuda" if torch.cuda.is_available() else "cpu",
help="Device for model inference.",
)
parser.add_argument(
"--dtype",
type=str,
default="bfloat16" if torch.cuda.is_available() else "float32",
help="Torch dtype for model weights.",
)
args = parser.parse_args() args = parser.parse_args()
with torch.inference_mode(): with torch.inference_mode():
process_file(**vars(args)) main(
param_path=args.param_path,
input_path=args.input_path,
output_dir=args.output_dir,
text_key=args.text_key,
batch_size=args.batch_size,
max_length=args.max_length,
token_level=args.token_level,
max_samples=args.max_samples,
device=args.device,
dtype=args.dtype,
)
+153
View File
@@ -0,0 +1,153 @@
"""ROUGE evaluation (manual implementation, no external deps).
Computes ROUGE-1, ROUGE-2, ROUGE-L precision, recall, and F1.
Usage::
# Batch evaluation from JSONL (each line: {"reference": ..., "candidate": ...})
python scripts/eval/evaluate_rouge.py --data_path preds.jsonl --output results.json
# As a library
from scripts.eval.evaluate_rouge import compute_rouge
scores = compute_rouge("the cat sat on the mat", "the cat sat")
"""
import argparse
import json
from collections import Counter
from typing import Dict, List, Tuple
def _tokenize(text: str) -> List[str]:
return text.split()
def _ngrams(tokens: List[str], n: int) -> Counter:
return Counter(zip(*[tokens[i:] for i in range(n)]))
def _lcs(x: List[str], y: List[str]) -> int:
m, n = len(x), len(y)
dp = [[0] * (n + 1) for _ in range(m + 1)]
for i in range(1, m + 1):
xi = x[i - 1]
dpi = dp[i]
dpi_1 = dp[i - 1]
for j in range(1, n + 1):
if xi == y[j - 1]:
dpi[j] = dpi_1[j - 1] + 1
else:
dpi[j] = dpi_1[j] if dpi_1[j] > dpi[j - 1] else dpi[j - 1]
return dp[m][n]
def _f1(precision: float, recall: float) -> float:
if precision + recall == 0:
return 0.0
return 2 * precision * recall / (precision + recall)
def _rouge_n(ref_tokens: List[str], cand_tokens: List[str], n: int) -> Dict[str, float]:
ref_ngrams = _ngrams(ref_tokens, n)
cand_ngrams = _ngrams(cand_tokens, n)
overlap = sum((cand_ngrams & ref_ngrams).values())
cand_total = sum(cand_ngrams.values())
ref_total = sum(ref_ngrams.values())
precision = overlap / cand_total if cand_total > 0 else 0.0
recall = overlap / ref_total if ref_total > 0 else 0.0
f1 = _f1(precision, recall)
return {"precision": precision, "recall": recall, "f1": f1}
def _rouge_l(ref_tokens: List[str], cand_tokens: List[str]) -> Dict[str, float]:
lcs_len = _lcs(ref_tokens, cand_tokens)
ref_len = len(ref_tokens)
cand_len = len(cand_tokens)
recall = lcs_len / ref_len if ref_len > 0 else 0.0
precision = lcs_len / cand_len if cand_len > 0 else 0.0
f1 = _f1(precision, recall)
return {"precision": precision, "recall": recall, "f1": f1}
def compute_rouge(
reference: str, candidate: str, n: int = 2
) -> Dict[str, Dict[str, float]]:
"""Compute ROUGE-N (1..n) and ROUGE-L scores.
Returns::
{
"rouge-1": {"precision": ..., "recall": ..., "f1": ...},
"rouge-2": {"precision": ..., "recall": ..., "f1": ...},
"rouge-l": {"precision": ..., "recall": ..., "f1": ...},
}
"""
ref_tokens = _tokenize(reference)
cand_tokens = _tokenize(candidate)
results = {}
for i in range(1, n + 1):
results[f"rouge-{i}"] = _rouge_n(ref_tokens, cand_tokens, i)
results["rouge-l"] = _rouge_l(ref_tokens, cand_tokens)
return results
def evaluate_file(data_path: str) -> Dict:
with open(data_path, "r", encoding="utf-8") as f:
pairs = [json.loads(line) for line in f if line.strip()]
agg = {
k: {"precision": 0.0, "recall": 0.0, "f1": 0.0}
for k in ("rouge-1", "rouge-2", "rouge-l")
}
per_item = []
for item in pairs:
ref = item["reference"]
cand = item["candidate"]
scores = compute_rouge(ref, cand)
per_item.append({**item, "scores": scores})
for k, v in scores.items():
agg[k]["precision"] += v["precision"]
agg[k]["recall"] += v["recall"]
agg[k]["f1"] += v["f1"]
n = len(pairs)
for k in agg:
agg[k] = {m: v / n for m, v in agg[k].items()}
return {"num_samples": n, "aggregate": agg, "per_item": per_item}
def main():
parser = argparse.ArgumentParser(description="ROUGE evaluation")
parser.add_argument(
"--data_path", required=True, help="JSONL with reference/candidate per line"
)
parser.add_argument("--output", type=str, default=None, help="Output JSON path")
args = parser.parse_args()
results = evaluate_file(args.data_path)
agg = results["aggregate"]
print(f"Samples: {results['num_samples']}")
print()
for metric in ("rouge-1", "rouge-2", "rouge-l"):
s = agg[metric]
print(
f" {metric:8s} P={s['precision']:.4f} R={s['recall']:.4f} F1={s['f1']:.4f}"
)
if args.output:
with open(args.output, "w", encoding="utf-8") as f:
json.dump(results, f, indent=2, ensure_ascii=False)
print(f"\nSaved to {args.output}")
if __name__ == "__main__":
main()
+123 -73
View File
@@ -1,12 +1,13 @@
"""Benchmark AutoRegressiveLM with KVCache""" """Benchmark AutoRegressiveLM with KVCache"""
import argparse
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, Dict from typing import Any, Dict
import torch import torch
from astrai.config import AutoRegressiveLMConfig from astrai.config import AutoRegressiveLMConfig
from astrai.inference import KVCache from astrai.inference import ContiguousCache, PageCache
from astrai.model.transformer import AutoRegressiveLM from astrai.model.transformer import AutoRegressiveLM
@@ -24,41 +25,14 @@ class GenerationBenchmark:
config: AutoRegressiveLMConfig, config: AutoRegressiveLMConfig,
device: str = "cuda", device: str = "cuda",
dtype: torch.dtype = torch.bfloat16, dtype: torch.dtype = torch.bfloat16,
page_size: int = 128, cache_type: str = "contiguous",
): ):
self.config = config self.config = config
self.device = device self.device = device
self.dtype = dtype self.dtype = dtype
self.cache_type = cache_type
self.model = AutoRegressiveLM(config).to(device=device, dtype=dtype) self.model = AutoRegressiveLM(config).to(device=device, dtype=dtype)
self.model.eval() self.model.eval()
head_dim = config.dim // config.n_heads
n_pages = (config.max_len * 4 + page_size - 1) // page_size
self._page_cache = KVCache(
config.n_layers,
n_pages,
page_size,
config.n_kv_heads,
head_dim,
device,
dtype,
)
def _prepare_inputs(self, batch_size: int, prompt_length: int, total_length: int):
prompt_ids = torch.randint(
low=0,
high=self.config.vocab_size,
size=(batch_size, prompt_length),
device=self.device,
dtype=torch.long,
)
gen_ids = torch.randint(
low=0,
high=self.config.vocab_size,
size=(batch_size, total_length - prompt_length),
device=self.device,
dtype=torch.long,
)
return prompt_ids, gen_ids
@torch.inference_mode() @torch.inference_mode()
def run_prefill_benchmark( def run_prefill_benchmark(
@@ -68,8 +42,12 @@ class GenerationBenchmark:
num_trials: int = 10, num_trials: int = 10,
) -> BenchmarkResult: ) -> BenchmarkResult:
for _ in range(3): for _ in range(3):
prompt_ids, _ = self._prepare_inputs( prompt_ids = torch.randint(
batch_size, prompt_length, prompt_length 0,
self.config.vocab_size,
(batch_size, prompt_length),
device=self.device,
dtype=torch.long,
) )
_ = self.model(prompt_ids) _ = self.model(prompt_ids)
torch.cuda.synchronize() torch.cuda.synchronize()
@@ -78,12 +56,15 @@ class GenerationBenchmark:
total_tokens = batch_size * prompt_length * num_trials total_tokens = batch_size * prompt_length * num_trials
for trial in range(num_trials): for trial in range(num_trials):
prompt_ids, _ = self._prepare_inputs( prompt_ids = torch.randint(
batch_size, prompt_length, prompt_length 0,
self.config.vocab_size,
(batch_size, prompt_length),
device=self.device,
dtype=torch.long,
) )
start = torch.cuda.Event(enable_timing=True) start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True) end = torch.cuda.Event(enable_timing=True)
start.record() start.record()
_ = self.model(prompt_ids) _ = self.model(prompt_ids)
end.record() end.record()
@@ -107,6 +88,7 @@ class GenerationBenchmark:
"prompt_length": prompt_length, "prompt_length": prompt_length,
"dtype": str(self.dtype), "dtype": str(self.dtype),
"device": self.device, "device": self.device,
"cache": "none",
}, },
) )
@@ -120,29 +102,56 @@ class GenerationBenchmark:
) -> BenchmarkResult: ) -> BenchmarkResult:
total_time = 0.0 total_time = 0.0
total_tokens = batch_size * gen_length * num_trials total_tokens = batch_size * gen_length * num_trials
page_size = self._page_cache.page_size
for trial in range(num_trials): for trial in range(num_trials):
prompt_ids, gen_ids = self._prepare_inputs( prompt_ids = torch.randint(
batch_size, 0,
prompt_length, self.config.vocab_size,
prompt_length + gen_length, (batch_size, prompt_length),
)
n_pages = (prompt_length + gen_length + page_size - 1) // page_size
total = n_pages * batch_size
pages = []
for _ in range(total):
p = self._page_cache._pool.alloc()
assert p >= 0, "OOM"
pages.append(p)
page_table = torch.tensor(
[pages[i * n_pages : (i + 1) * n_pages] for i in range(batch_size)],
dtype=torch.long,
device=self.device, device=self.device,
dtype=torch.long,
)
gen_ids = torch.randint(
0,
self.config.vocab_size,
(batch_size, gen_length),
device=self.device,
dtype=torch.long,
) )
cv = self._page_cache.bind(page_table, total_len=prompt_length) head_dim = self.config.dim // self.config.n_heads
max_seq = prompt_length + gen_length
if self.cache_type == "contiguous":
cache = ContiguousCache(
self.config.n_layers,
batch_size,
max_seq,
self.config.n_kv_heads,
head_dim,
self.device,
self.dtype,
)
else:
page_size = 128
n_pages = (max_seq + page_size - 1) // page_size * batch_size
cache = PageCache(
self.config.n_layers,
n_pages,
page_size,
self.config.n_kv_heads,
head_dim,
self.device,
self.dtype,
)
task_ids = [f"b{i}" for i in range(batch_size)]
for tid in task_ids:
cache.task_alloc(tid, [0] * max_seq)
for p in range(max_seq):
cache.task_extend(tid, p)
cv = cache.bind_tasks(task_ids, prompt_length, self.device)
_ = self.model( _ = self.model(
prompt_ids, prompt_ids,
paged_cache=cv, paged_cache=cv,
@@ -152,37 +161,35 @@ class GenerationBenchmark:
.unsqueeze(0) .unsqueeze(0)
.expand(batch_size, -1), .expand(batch_size, -1),
) )
torch.cuda.synchronize() torch.cuda.synchronize()
start = torch.cuda.Event(enable_timing=True) start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True) end = torch.cuda.Event(enable_timing=True)
start.record() start.record()
current_pos = prompt_length
for i in range(gen_length): for i in range(gen_length):
input_token = gen_ids[:, i : i + 1] pos = prompt_length + i
cv = self._page_cache.bind(page_table, total_len=current_pos + 1) cv = cache.bind_tasks(task_ids, pos + 1, self.device)
_ = self.model( _ = self.model(
input_token, gen_ids[:, i : i + 1],
paged_cache=cv, paged_cache=cv,
position_ids=torch.full( position_ids=torch.full(
(batch_size, 1), (batch_size, 1),
current_pos, pos,
dtype=torch.long, dtype=torch.long,
device=self.device, device=self.device,
), ),
) )
current_pos += 1
end.record() end.record()
torch.cuda.synchronize() torch.cuda.synchronize()
for tid in task_ids:
cache.task_free(tid)
trial_time = start.elapsed_time(end) / 1000 trial_time = start.elapsed_time(end) / 1000
total_time += trial_time total_time += trial_time
for idx in pages:
self._page_cache._pool.free(idx)
print( print(
f" Trial {trial + 1}/{num_trials}: {gen_length} tokens in {trial_time:.3f}s " f" Trial {trial + 1}/{num_trials}: {gen_length} tokens in {trial_time:.3f}s "
f"({gen_length / trial_time:.1f} tok/s)" f"({gen_length / trial_time:.1f} tok/s)"
@@ -199,6 +206,7 @@ class GenerationBenchmark:
"gen_length": gen_length, "gen_length": gen_length,
"dtype": str(self.dtype), "dtype": str(self.dtype),
"device": self.device, "device": self.device,
"cache": self.cache_type,
}, },
) )
@@ -216,6 +224,42 @@ def print_benchmark_result(result: BenchmarkResult):
if __name__ == "__main__": if __name__ == "__main__":
parser = argparse.ArgumentParser(description="AutoRegressiveLM benchmark")
parser.add_argument(
"--device", type=str, default="cuda", help="Device (default: cuda)"
)
parser.add_argument(
"--dtype",
type=str,
default="bfloat16",
choices=["bfloat16", "float16", "float32"],
help="Dtype",
)
parser.add_argument(
"--cache",
type=str,
default="contiguous",
choices=["contiguous", "paged"],
help="KV cache type",
)
parser.add_argument("--batch_size", type=int, default=4, help="Batch size")
parser.add_argument("--prompt_length", type=int, default=512, help="Prompt length")
parser.add_argument("--gen_length", type=int, default=128, help="Generation length")
parser.add_argument("--num_trials", type=int, default=5, help="Number of trials")
parser.add_argument(
"--prefill_only", action="store_true", help="Run prefill benchmark only"
)
parser.add_argument(
"--decode_only", action="store_true", help="Run decoding benchmark only"
)
args = parser.parse_args()
dtype_map = {
"bfloat16": torch.bfloat16,
"float16": torch.float16,
"float32": torch.float32,
}
config = AutoRegressiveLMConfig( config = AutoRegressiveLMConfig(
vocab_size=10000, vocab_size=10000,
dim=1536, dim=1536,
@@ -227,23 +271,29 @@ if __name__ == "__main__":
norm_eps=1e-5, norm_eps=1e-5,
) )
benchmark = GenerationBenchmark(config) benchmark = GenerationBenchmark(
config, device=args.device, dtype=dtype_map[args.dtype], cache_type=args.cache
)
print("=" * 80) print("=" * 80)
print("Running AutoRegressiveLM Generation Benchmark (KVCache)") print(
f"Running AutoRegressiveLM Benchmark (device={args.device}, dtype={args.dtype})"
)
print("=" * 80) print("=" * 80)
if not args.decode_only:
prefill_result = benchmark.run_prefill_benchmark( prefill_result = benchmark.run_prefill_benchmark(
batch_size=4, batch_size=args.batch_size,
prompt_length=512, prompt_length=args.prompt_length,
num_trials=5, num_trials=args.num_trials,
) )
print_benchmark_result(prefill_result) print_benchmark_result(prefill_result)
if not args.prefill_only:
gen_result = benchmark.run_decoding_benchmark( gen_result = benchmark.run_decoding_benchmark(
batch_size=4, batch_size=args.batch_size,
prompt_length=512, prompt_length=args.prompt_length,
gen_length=128, gen_length=args.gen_length,
num_trials=5, num_trials=args.num_trials,
) )
print_benchmark_result(gen_result) print_benchmark_result(gen_result)
+105 -25
View File
@@ -1,7 +1,10 @@
import argparse import argparse
import json import json
import time
from typing import Optional
import torch import torch
from tqdm import tqdm
from astrai.inference import InferenceEngine from astrai.inference import InferenceEngine
from astrai.model import AutoModel from astrai.model import AutoModel
@@ -17,62 +20,109 @@ def processor(
top_p: float, top_p: float,
question_key: str, question_key: str,
response_key: str, response_key: str,
max_tokens: int, max_tokens: Optional[int],
batch_size: int, batch_size: int,
num_samples: int = 1,
cache_len: int = 2048,
frequency_penalty: float = 0.0,
rep_window: int = 64,
): ):
# Load model and tokenizer print(f"Loading model from {param_path} ...")
t0 = time.time()
model = AutoModel.from_pretrained(param_path) model = AutoModel.from_pretrained(param_path)
tokenizer = AutoTokenizer.from_pretrained(param_path) tokenizer = AutoTokenizer.from_pretrained(param_path)
model.to(device="cuda", dtype=torch.bfloat16) model.to(device="cuda", dtype=torch.bfloat16)
print(f" model loaded in {time.time() - t0:.1f}s")
# Create inference engine
engine = InferenceEngine( engine = InferenceEngine(
model=model, tokenizer=tokenizer, max_batch_size=batch_size model=model,
tokenizer=tokenizer,
max_batch_size=batch_size * num_samples,
max_seq_len=cache_len,
max_prompt_len=cache_len,
) )
print(f"Reading {input_json_file} ...")
with open(input_json_file, "r", encoding="utf-8") as f: with open(input_json_file, "r", encoding="utf-8") as f:
input_data = [json.loads(line) for line in f] input_data = [json.loads(line) for line in f]
# Check input format: chat messages or raw text
if input_data and "messages" in input_data[0]: if input_data and "messages" in input_data[0]:
# Chat format: [{"messages": [...]}]
prompts = [ prompts = [
tokenizer.apply_chat_template(item["messages"], tokenize=False) tokenizer.apply_chat_template(item["messages"], tokenize=False)
for item in input_data for item in input_data
] ]
else: else:
# Raw text format: [{"question": "..."}]
prompts = [item[question_key] for item in input_data] prompts = [item[question_key] for item in input_data]
print(f" {len(prompts)} prompts loaded\n")
# Use provided max_tokens or default to model config max_len
if max_tokens is None: if max_tokens is None:
max_tokens = model.config.max_len max_tokens = model.config.max_len
# Generate responses (batch) chunk_size = max(1, batch_size)
responses = engine.generate(
prompt=prompts, with open(output_json_file, "w", encoding="utf-8") as f:
pbar = tqdm(
total=len(prompts) * num_samples,
unit="gen",
desc=f" Generating ({num_samples}x/prompt)",
)
for chunk_start in range(0, len(prompts), chunk_size):
chunk = prompts[chunk_start : chunk_start + chunk_size]
if num_samples > 1:
chunk_expanded = [p for p in chunk for _ in range(num_samples)]
resp_chunk = engine.generate(
prompt=chunk_expanded,
stream=False, stream=False,
max_tokens=max_tokens, max_tokens=max_tokens,
temperature=temperature, temperature=temperature,
top_p=top_p, top_p=top_p,
top_k=top_k, top_k=top_k,
frequency_penalty=frequency_penalty,
rep_window=rep_window,
)
resp_chunk = [
resp_chunk[i * num_samples : (i + 1) * num_samples]
for i in range(len(chunk))
]
else:
resp_chunk = engine.generate(
prompt=chunk,
stream=False,
max_tokens=max_tokens,
temperature=temperature,
top_p=top_p,
top_k=top_k,
frequency_penalty=frequency_penalty,
rep_window=rep_window,
) )
# Write results for i, prompt in enumerate(chunk):
with open(output_json_file, "w", encoding="utf-8") as f:
for prompt, response in zip(prompts, responses):
if input_data and "messages" in input_data[0]: if input_data and "messages" in input_data[0]:
output_item = {"response": response} orig = input_data[chunk_start + i]
output_item = {**orig, response_key: resp_chunk[i]}
else: else:
output_item = {question_key: prompt, response_key: response} output_item = {
question_key: prompt,
response_key: resp_chunk[i],
}
f.write(json.dumps(output_item, ensure_ascii=False) + "\n") f.write(json.dumps(output_item, ensure_ascii=False) + "\n")
# Cleanup pbar.update(len(chunk) * num_samples)
pbar.close()
elapsed = time.time() - t0
print(
f"\nDone! {len(prompts)} prompts x {num_samples} samples -> {output_json_file}"
)
print(f"Total time: {elapsed:.1f}s ({elapsed / len(prompts):.2f}s/prompt)")
engine.shutdown() engine.shutdown()
if __name__ == "__main__": if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Run generate with a Khaosz model.") parser = argparse.ArgumentParser(description="Batch generation from JSONL file.")
parser.add_argument( parser.add_argument(
"--param_path", type=str, required=True, help="Path to the model directory." "--param_path", type=str, required=True, help="Path to the model directory."
@@ -93,38 +143,68 @@ if __name__ == "__main__":
"--question_key", "--question_key",
type=str, type=str,
default="question", default="question",
help="Key for the question in the input JSON.", help="Key for the question in the input JSON (default: question).",
) )
parser.add_argument( parser.add_argument(
"--response_key", "--response_key",
type=str, type=str,
default="response", default="response",
help="Key for the response in the output JSON.", help="Key for the response in the output JSON (default: response).",
) )
parser.add_argument( parser.add_argument(
"--temperature", "--temperature",
type=float, type=float,
default=0.60, default=0.60,
help="Temperature for generating responses.", help="Temperature for generating responses (default: 0.60).",
) )
parser.add_argument( parser.add_argument(
"--top_k", type=int, default=30, help="Top-k value for generating responses." "--top_k",
type=int,
default=30,
help="Top-k value for generating responses (default: 30).",
) )
parser.add_argument( parser.add_argument(
"--top_p", "--top_p",
type=float, type=float,
default=0.95, default=0.95,
help="Top-p value for generating responses.", help="Top-p value for generating responses (default: 0.95).",
) )
parser.add_argument( parser.add_argument(
"--batch_size", type=int, default=1, help="Batch size for generating responses." "--batch_size",
type=int,
default=1,
help="Batch size for generating responses (default: 1).",
)
parser.add_argument(
"--num_samples",
type=int,
default=1,
help="Number of responses per prompt (expands batch internally, default: 1).",
) )
parser.add_argument( parser.add_argument(
"--max_tokens", "--max_tokens",
type=int, type=int,
default=2048, default=None,
help="Maximum tokens to generate (default: model config max_len).", help="Maximum tokens to generate (default: model config max_len).",
) )
parser.add_argument(
"--cache_len",
type=int,
default=2048,
help="KV cache & prompt truncation length (default: 2048, lower = less memory).",
)
parser.add_argument(
"--frequency_penalty",
type=float,
default=0.0,
help="Frequency penalty to reduce repetition (default: 0.0, try 0.5-1.0).",
)
parser.add_argument(
"--rep_window",
type=int,
default=64,
help="Window size for frequency penalty (default: 64).",
)
args = parser.parse_args() args = parser.parse_args()
+211 -60
View File
@@ -1,17 +1,98 @@
import argparse import argparse
import os import os
from functools import partial from functools import partial
from typing import Any, Dict
import torch import torch
import torch.optim as optim import torch.optim as optim
from torch import Tensor, nn
from astrai.config import AutoRegressiveLMConfig, TrainConfig from astrai.config import AutoRegressiveLMConfig, TrainConfig
from astrai.dataset import DatasetFactory from astrai.dataset import DatasetFactory, dpo_collate_fn, grpo_collate_fn
from astrai.model import AutoRegressiveLM from astrai.model import AutoRegressiveLM
from astrai.model.components.decoder_block import DecoderBlock from astrai.model.components.decoder_block import DecoderBlock
from astrai.trainer import SchedulerFactory, Trainer from astrai.trainer import SchedulerFactory, Trainer
class MuonMix(optim.Optimizer):
"""Combined Muon (matrix) + AdamW (non-matrix) optimizer."""
def __init__(
self,
model: nn.Module,
lr: float = 3e-4,
weight_decay: float = 0.1,
momentum: float = 0.95,
nesterov: bool = True,
ns_steps: int = 5,
adjust_lr_fn: str = "match_rms_adamw",
):
defaults = dict(
lr=lr,
weight_decay=weight_decay,
momentum=momentum,
nesterov=nesterov,
ns_steps=ns_steps,
adjust_lr_fn=adjust_lr_fn,
)
params = [p for p in model.parameters() if p.requires_grad]
super().__init__(params, defaults)
matrix_params: list[Tensor] = []
other_params: list[Tensor] = []
for name, param in model.named_parameters():
if not param.requires_grad:
continue
if (
param.dim() >= 2
and "norm" not in name
and "bias" not in name
and "embed" not in name
and "lm_head" not in name
):
matrix_params.append(param)
else:
other_params.append(param)
self.muon = optim.Muon(
matrix_params,
lr=lr,
weight_decay=weight_decay,
momentum=momentum,
nesterov=nesterov,
ns_steps=ns_steps,
adjust_lr_fn=adjust_lr_fn,
)
self.adamw = optim.AdamW(
[{"params": other_params, "weight_decay": 0.0}],
lr=lr,
betas=(0.9, 0.95),
fused=True,
)
self.param_groups = [*self.muon.param_groups, *self.adamw.param_groups]
@torch.no_grad()
def step(self, closure=None):
self.muon.step(closure)
self.adamw.step(closure)
def zero_grad(self, set_to_none: bool = True):
self.muon.zero_grad(set_to_none)
self.adamw.zero_grad(set_to_none)
def state_dict(self) -> Dict[str, Any]:
return {
"muon": self.muon.state_dict(),
"adamw": self.adamw.state_dict(),
}
def load_state_dict(self, state_dict: Dict[str, Any]):
self.muon.load_state_dict(state_dict["muon"])
self.adamw.load_state_dict(state_dict["adamw"])
self.param_groups = [*self.muon.param_groups, *self.adamw.param_groups]
def parse_args() -> argparse.Namespace: def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Train the AutoRegressiveLM model.") parser = argparse.ArgumentParser(description="Train the AutoRegressiveLM model.")
@@ -35,6 +116,13 @@ def parse_args() -> argparse.Namespace:
required=True, required=True,
help="Path to the model parameters or resume checkpoint.", help="Path to the model parameters or resume checkpoint.",
) )
parser.add_argument(
"--resume",
action="store_true",
default=False,
help="Resume training from checkpoint at --param_path "
"(restore epoch, consumed_samples, optimizer & scheduler state).",
)
parser.add_argument( parser.add_argument(
"--n_epoch", type=int, default=1, help="Number of epochs to train." "--n_epoch", type=int, default=1, help="Number of epochs to train."
@@ -60,26 +148,39 @@ def parse_args() -> argparse.Namespace:
parser.add_argument( parser.add_argument(
"--max_grad_norm", "--max_grad_norm",
type=float, type=float,
default=1.0, default=None,
help="Max gradient norm for clipping.", help="Max gradient norm for clipping. None disables clipping.",
) )
parser.add_argument( parser.add_argument(
"--adamw_beta1", "--weight_decay",
type=float, type=float,
default=0.9, default=0.1,
help="Beta1 for AdamW optimizer.", help="Weight decay (applied to Muon matrix params; non-matrix use 0).",
) )
parser.add_argument( parser.add_argument(
"--adamw_beta2", "--muon_momentum",
type=float, type=float,
default=0.95, default=0.95,
help="Beta2 for AdamW optimizer.", help="Momentum factor for Muon optimizer.",
) )
parser.add_argument( parser.add_argument(
"--adamw_weight_decay", "--muon_nesterov",
type=float, action=argparse.BooleanOptionalAction,
default=0.01, default=True,
help="Weight decay for AdamW optimizer.", help="Enable Nesterov momentum for Muon.",
)
parser.add_argument(
"--muon_ns_steps",
type=int,
default=5,
help="Newton-Schulz iteration steps for Muon.",
)
parser.add_argument(
"--muon_adjust_lr",
type=str,
default="match_rms_adamw",
choices=["original", "match_rms_adamw"],
help="Muon learning rate adjustment strategy.",
) )
parser.add_argument( parser.add_argument(
"--random_seed", type=int, default=3407, help="Random seed for reproducibility." "--random_seed", type=int, default=3407, help="Random seed for reproducibility."
@@ -150,8 +251,8 @@ def parse_args() -> argparse.Namespace:
parser.add_argument( parser.add_argument(
"--metrics", "--metrics",
nargs="*", nargs="*",
default=["loss", "lr"], default=["loss", "lr", "grad_norm"],
help="Metrics to log (e.g. --metrics loss lr val_loss). Default: loss lr.", help="Metrics to log (e.g. --metrics loss lr val_loss). Default: loss lr grad_norm.",
) )
parser.add_argument( parser.add_argument(
"--log_dir", "--log_dir",
@@ -159,23 +260,14 @@ def parse_args() -> argparse.Namespace:
default="checkpoint/logs", default="checkpoint/logs",
help="Directory for metric logs.", help="Directory for metric logs.",
) )
parser.add_argument(
"--log_interval",
type=int,
default=100,
help="Number of batch iterations between metric logs.",
)
parser.add_argument(
"--grpo_sync_interval",
type=int,
default=200,
help="GRPO ref model sync interval (steps).",
)
parser.add_argument( parser.add_argument(
"--start_epoch", type=int, default=0, help="Start epoch for training." "--start_epoch", type=int, default=0, help="Start epoch for training."
) )
parser.add_argument( parser.add_argument(
"--start_batch", type=int, default=0, help="Start batch for training." "--start_samples",
type=int,
default=0,
help="Start samples (per rank) for training.",
) )
parser.add_argument( parser.add_argument(
@@ -221,6 +313,44 @@ def parse_args() -> argparse.Namespace:
help="NEFTune noise alpha (0=disabled, typical: 5.0).", help="NEFTune noise alpha (0=disabled, typical: 5.0).",
) )
parser.add_argument(
"--schedule_type",
type=str,
default="cosine",
choices=["cosine", "sgdr", "wsd"],
help="Learning rate scheduler type.",
)
parser.add_argument(
"--min_rate",
type=float,
default=None,
help="Minimum LR as fraction of base LR. Uses scheduler default if not set (cosine/sgdr: 0.05, wsd: 0.0).",
)
parser.add_argument(
"--cycle_length",
type=int,
default=None,
help="SGDR first cycle length in steps. Defaults to total_steps - warmup_steps.",
)
parser.add_argument(
"--t_mult",
type=int,
default=2,
help="SGDR cycle length multiplier per restart.",
)
parser.add_argument(
"--stable_steps",
type=int,
default=None,
help="WSD stable plateau steps. Required when --schedule_type wsd.",
)
parser.add_argument(
"--decay_steps",
type=int,
default=None,
help="WSD decay steps. Defaults to total_steps - warmup_steps - stable_steps.",
)
args = parser.parse_args() args = parser.parse_args()
return args return args
@@ -230,8 +360,8 @@ def create_model(config):
return AutoRegressiveLM(config).to(dtype=torch.bfloat16) return AutoRegressiveLM(config).to(dtype=torch.bfloat16)
def create_optimizer(model, **kwargs) -> optim.Optimizer: def create_optimizer(model, **kwargs) -> MuonMix:
return optim.AdamW(model.parameters(), fused=True, **kwargs) return MuonMix(model, **kwargs)
def create_scheduler( def create_scheduler(
@@ -262,11 +392,11 @@ def train(
train_type: str, train_type: str,
param_path: str, param_path: str,
data_root_path: str, data_root_path: str,
max_lr: float, resume: bool,
n_epoch: int, n_epoch: int,
batch_per_device: int, batch_per_device: int,
start_epoch: int, start_epoch: int,
start_batch: int, start_samples: int,
grad_accum_steps: int, grad_accum_steps: int,
warmup_ratio: float, warmup_ratio: float,
ckpt_interval: int, ckpt_interval: int,
@@ -275,17 +405,7 @@ def train(
val_step: int, val_step: int,
metrics: list[str], metrics: list[str],
log_dir: str, log_dir: str,
log_interval: int,
dpo_beta: float,
grpo_clip_eps: float,
grpo_kl_coef: float,
group_size: int,
grpo_sync_interval: int,
adamw_beta1: float,
adamw_beta2: float,
adamw_weight_decay: float,
max_grad_norm: float, max_grad_norm: float,
label_smoothing: float,
random_seed: int, random_seed: int,
num_workers: int, num_workers: int,
pin_memory: bool, pin_memory: bool,
@@ -300,6 +420,13 @@ def train(
master_port: str, master_port: str,
start_method: str, start_method: str,
neftune_alpha: float, neftune_alpha: float,
schedule_type: str,
min_rate: float,
cycle_length: int,
t_mult: int,
stable_steps: int,
decay_steps: int,
**kwargs,
): ):
assert train_type in ["seq", "sft", "dpo", "grpo"] assert train_type in ["seq", "sft", "dpo", "grpo"]
assert os.path.exists(param_path) assert os.path.exists(param_path)
@@ -309,17 +436,17 @@ def train(
# Load config # Load config
config_path = os.path.join(param_path, "config.json") config_path = os.path.join(param_path, "config.json")
config = AutoRegressiveLMConfig.from_file(config_path) config = AutoRegressiveLMConfig.from_file(config_path)
config.neftune_alpha = neftune_alpha
if window_size is None: if window_size is None:
window_size = config.max_len window_size = config.max_len
strategy_kwargs = { strategy_kwargs = {
"beta": dpo_beta, "beta": kwargs.pop("dpo_beta"),
"label_smoothing": label_smoothing, "label_smoothing": kwargs.pop("label_smoothing"),
"clip_eps": grpo_clip_eps, "clip_eps": kwargs.pop("grpo_clip_eps"),
"kl_coef": grpo_kl_coef, "kl_coef": kwargs.pop("grpo_kl_coef"),
"group_size": group_size, "group_size": kwargs.pop("group_size"),
"sync_interval": grpo_sync_interval,
} }
executor_kwargs = { executor_kwargs = {
@@ -333,33 +460,57 @@ def train(
load_path=data_root_path, load_path=data_root_path,
window_size=window_size, window_size=window_size,
stride=stride, stride=stride,
tokenizer_path=param_path,
) )
optimizer_fn = partial( optimizer_fn = partial(
create_optimizer, create_optimizer,
**{ lr=kwargs.pop("max_lr"),
"lr": max_lr, weight_decay=kwargs.pop("weight_decay"),
"betas": (adamw_beta1, adamw_beta2), momentum=kwargs.pop("muon_momentum"),
"weight_decay": adamw_weight_decay, nesterov=kwargs.pop("muon_nesterov"),
}, ns_steps=kwargs.pop("muon_ns_steps"),
adjust_lr_fn=kwargs.pop("muon_adjust_lr"),
) )
total_steps = compute_total_steps( total_steps = compute_total_steps(
len(dataset), n_epoch, batch_per_device, nprocs, grad_accum_steps len(dataset), n_epoch, batch_per_device, nprocs, grad_accum_steps
) )
warmup_steps = int(warmup_ratio * total_steps) warmup_steps = int(warmup_ratio * total_steps)
warmup_steps = min(warmup_steps, total_steps)
scheduler_kwargs = {"warmup_steps": warmup_steps}
if schedule_type == "cosine":
scheduler_kwargs["lr_decay_steps"] = total_steps - warmup_steps
elif schedule_type == "sgdr":
scheduler_kwargs["cycle_length"] = cycle_length or (total_steps - warmup_steps)
scheduler_kwargs["t_mult"] = t_mult
elif schedule_type == "wsd":
remaining = total_steps - warmup_steps
stable_steps_ = stable_steps or max(1, int(remaining * 0.8))
scheduler_kwargs["stable_steps"] = stable_steps_
scheduler_kwargs["decay_steps"] = max(
1, decay_steps or (remaining - stable_steps_)
)
if min_rate is not None:
scheduler_kwargs["min_rate"] = min_rate
scheduler_fn = partial( scheduler_fn = partial(
create_scheduler, create_scheduler,
**{ schedule_type=schedule_type,
"schedule_type": "cosine", **scheduler_kwargs,
"warmup_steps": min(warmup_steps, total_steps),
"lr_decay_steps": total_steps - min(warmup_steps, total_steps),
},
) )
grad_ckpt_modules = [DecoderBlock] if gradient_checkpointing else [] grad_ckpt_modules = [DecoderBlock] if gradient_checkpointing else []
collate_fn = None
if train_type == "dpo":
collate_fn = dpo_collate_fn
elif train_type == "grpo":
collate_fn = grpo_collate_fn
train_config = TrainConfig( train_config = TrainConfig(
model_fn=model_fn, model_fn=model_fn,
strategy=train_type, strategy=train_type,
@@ -370,7 +521,7 @@ def train(
n_epoch=n_epoch, n_epoch=n_epoch,
batch_per_device=batch_per_device, batch_per_device=batch_per_device,
start_epoch=start_epoch, start_epoch=start_epoch,
start_batch=start_batch, start_samples=start_samples,
ckpt_interval=ckpt_interval, ckpt_interval=ckpt_interval,
grad_accum_steps=grad_accum_steps, grad_accum_steps=grad_accum_steps,
max_grad_norm=max_grad_norm, max_grad_norm=max_grad_norm,
@@ -388,15 +539,15 @@ def train(
val_step=val_step, val_step=val_step,
metrics=metrics, metrics=metrics,
log_dir=log_dir, log_dir=log_dir,
log_interval=log_interval,
gradient_checkpointing_modules=grad_ckpt_modules, gradient_checkpointing_modules=grad_ckpt_modules,
executor_kwargs=executor_kwargs, executor_kwargs=executor_kwargs,
extra_kwargs=strategy_kwargs, extra_kwargs=strategy_kwargs,
neftune_alpha=neftune_alpha, neftune_alpha=neftune_alpha,
collate_fn=collate_fn,
) )
trainer = Trainer(train_config) trainer = Trainer(train_config)
trainer.train(resume_dir=param_path) trainer.train(param_path=param_path, resume=resume)
if __name__ == "__main__": if __name__ == "__main__":
+61
View File
@@ -0,0 +1,61 @@
import os
import sys
from pathlib import Path
from setuptools import setup
from setuptools.command.build_ext import build_ext as _build_ext
sys.path.insert(0, str(Path(__file__).parent))
os.makedirs("astrai/extension", exist_ok=True)
def _should_build():
force = os.environ.get("CSRC_KERNELS", "").strip().lower()
if force == "true":
return True
if force == "false":
return False
try:
import shutil
import torch
return shutil.which("nvcc") is not None and torch.cuda.is_available()
except Exception:
return False
ext_modules = []
cmdclass = {}
if _should_build():
import torch
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
from csrc.build import REGISTRY
_torch_lib = torch.utils.cpp_extension.library_paths()[0]
for name, info in REGISTRY.items():
ext_modules.append(
CUDAExtension(
f"astrai.extension.{name}",
info["sources"],
extra_compile_args={
"cxx": info["cxx_flags"],
"nvcc": info["nvcc_flags"],
},
extra_link_args=[f"-Wl,-rpath,{_torch_lib}"],
)
)
cmdclass["build_ext"] = BuildExtension
if not cmdclass:
class _NullBuildExt(_build_ext):
def build_extensions(self):
pass
cmdclass["build_ext"] = _NullBuildExt
setup(ext_modules=ext_modules, cmdclass=cmdclass)
+1 -1
View File
@@ -75,7 +75,7 @@ class MultiTurnDataset(Dataset):
class EarlyStoppingDataset(Dataset): class EarlyStoppingDataset(Dataset):
"""Dataset that triggers early stopping after a specified number of iterations.""" """Dataset that triggers early stopping after consuming a specified number of samples."""
def __init__(self, length=10, stop_after=5): def __init__(self, length=10, stop_after=5):
self.length = length self.length = length
+47
View File
@@ -1,3 +1,5 @@
import json
import os
import tempfile import tempfile
import pytest import pytest
@@ -8,6 +10,11 @@ from astrai.config.preprocess_config import (
PipelineConfig, PipelineConfig,
ProcessingConfig, ProcessingConfig,
) )
from astrai.preprocessing.builder import (
MultiOutputMaskBuilder,
SectionedMaskBuilder,
SingleOutputMaskBuilder,
)
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
_SPECIAL_TOKENS_CONFIG = { _SPECIAL_TOKENS_CONFIG = {
@@ -200,3 +207,43 @@ def make_grpo_no_template_config():
mask_default="mask", mask_default="mask",
preprocessing=ProcessingConfig(max_seq_len=2048), preprocessing=ProcessingConfig(max_seq_len=2048),
) )
@pytest.fixture
def builder():
return SectionedMaskBuilder()
@pytest.fixture
def single_builder():
return SingleOutputMaskBuilder()
@pytest.fixture
def multi_builder():
return MultiOutputMaskBuilder()
@pytest.fixture
def tokenizer_dir(temp_dir, test_tokenizer):
d = os.path.join(temp_dir, "tok")
os.makedirs(d, exist_ok=True)
test_tokenizer._tokenizer.save(os.path.join(d, "tokenizer.json"))
with open(os.path.join(d, "tokenizer_config.json"), "w") as f:
json.dump(
{"special_tokens": {"pad_token": "<|_pad_|>", "unk_token": "<|_unk_|>"}}, f
)
return d
@pytest.fixture
def chat_tokenizer_dir(temp_dir, chat_tokenizer):
d = os.path.join(temp_dir, "tok")
os.makedirs(d, exist_ok=True)
chat_tokenizer._tokenizer.save(os.path.join(d, "tokenizer.json"))
with open(os.path.join(d, "tokenizer_config.json"), "w") as f:
json.dump(
{"special_tokens": _SPECIAL_TOKENS_CONFIG, "chat_template": _CHAT_TEMPLATE},
f,
)
return d
+9 -4
View File
@@ -25,7 +25,9 @@ def test_single_process():
scheduler.step() scheduler.step()
checkpoint = Checkpoint(state_dict=model.state_dict(), epoch=3, iteration=30) checkpoint = Checkpoint(
state_dict=model.state_dict(), epoch=3, consumed_samples=120
)
with tempfile.TemporaryDirectory() as tmpdir: with tempfile.TemporaryDirectory() as tmpdir:
checkpoint.save(tmpdir) checkpoint.save(tmpdir)
@@ -33,7 +35,7 @@ def test_single_process():
loaded_checkpoint = Checkpoint.load(tmpdir) loaded_checkpoint = Checkpoint.load(tmpdir)
assert loaded_checkpoint.epoch == 3 assert loaded_checkpoint.epoch == 3
assert loaded_checkpoint.iteration == 30 assert loaded_checkpoint.consumed_samples == 120
def test_checkpoint_with_extra(): def test_checkpoint_with_extra():
@@ -46,7 +48,10 @@ def test_checkpoint_with_extra():
"scheduler": {"last_epoch": 5}, "scheduler": {"last_epoch": 5},
} }
checkpoint = Checkpoint( checkpoint = Checkpoint(
state_dict=model.state_dict(), epoch=1, iteration=10, extra=extra state_dict=model.state_dict(),
epoch=1,
consumed_samples=40,
extra=extra,
) )
with tempfile.TemporaryDirectory() as tmpdir: with tempfile.TemporaryDirectory() as tmpdir:
@@ -77,7 +82,7 @@ def simple_training():
checkpoint = Checkpoint( checkpoint = Checkpoint(
state_dict=model.state_dict(), state_dict=model.state_dict(),
epoch=2, epoch=2,
iteration=10, consumed_samples=40,
) )
rank = get_rank() rank = get_rank()
+816 -158
View File
File diff suppressed because it is too large Load Diff

Some files were not shown because too many files have changed in this diff Show More