docs: audit non-CUDA documentation
- Aligns CLI and strategy metric contracts - Refreshes architecture, dataflow, preprocessing, distributed, and eval guides - Corrects links, TOCs, defaults, and repository paths
This commit is contained in:
+25
-17
@@ -2,6 +2,18 @@
|
||||
|
||||
AstrAI supports three parallel modes: **single GPU** (`none`), **Data Parallel** (`ddp`), and **Fully Sharded Data Parallel** (`fsdp`). This guide covers when to use each, how to launch multi-GPU training, and how gradient accumulation works.
|
||||
|
||||
## Contents
|
||||
|
||||
- [Quick Start](#quick-start)
|
||||
- [Parallel Modes](#parallel-modes)
|
||||
- [Gradient Accumulation](#gradient-accumulation)
|
||||
- [Process Launching](#process-launching)
|
||||
- [NCCL Troubleshooting](#nccl-troubleshooting)
|
||||
- [Checkpoint Saving](#checkpoint-saving)
|
||||
- [Total Steps Calculation](#total-steps-calculation)
|
||||
- [Real Examples](#real-examples)
|
||||
- [CLI Parameters](#cli-parameters)
|
||||
|
||||
## Quick Start
|
||||
|
||||
### Single GPU
|
||||
@@ -21,9 +33,6 @@ python scripts/tools/train.py \
|
||||
|
||||
```bash
|
||||
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export NCCL_NET_GDR_LEVEL=0
|
||||
|
||||
python scripts/tools/train.py \
|
||||
--train_type=sft \
|
||||
--param_path ./params \
|
||||
@@ -38,9 +47,6 @@ python scripts/tools/train.py \
|
||||
|
||||
```bash
|
||||
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export NCCL_NET_GDR_LEVEL=0
|
||||
|
||||
python scripts/tools/train.py \
|
||||
--train_type=sft \
|
||||
--param_path ./params \
|
||||
@@ -110,7 +116,7 @@ AstrAI auto-detects the launch method:
|
||||
|
||||
| Detection | Strategy | Use Case |
|
||||
|-----------|----------|----------|
|
||||
| `torchelastic` / `torchrun` env vars | `TorchrunStrategy` | External orchestrator (torchrun, SLURM, K8s) |
|
||||
| `torchelastic` / `torchrun` env vars | `TorchrunStrategy` | External orchestrator (`torchrun`, K8s) |
|
||||
| `RANK` + `WORLD_SIZE` env vars | `TorchrunStrategy` | External launch |
|
||||
| Neither | `LocalStrategy` | `python scripts/tools/train.py` (in-process spawn) |
|
||||
|
||||
@@ -126,23 +132,28 @@ For multi-node or SLURM environments:
|
||||
torchrun --nproc_per_node=4 scripts/tools/train.py \
|
||||
--train_type=sft \
|
||||
--parallel_mode=ddp \
|
||||
--nprocs=4 \
|
||||
--param_path ./params \
|
||||
--data_root_path ./dataset \
|
||||
--batch_per_device=4
|
||||
```
|
||||
|
||||
When launched via torchrun, AstrAI reads `RANK`, `WORLD_SIZE`, `LOCAL_RANK` from the environment and uses `TorchrunStrategy`. The `--nprocs` flag is ignored (the orchestrator controls process count).
|
||||
When launched via `torchrun`, the launcher creates the worker processes. AstrAI reads `RANK`, `WORLD_SIZE`, and `LOCAL_RANK` from the environment and uses `TorchrunStrategy`; `--nprocs` does not control process creation in this mode.
|
||||
|
||||
## NCCL Environment Variables
|
||||
The current training CLI still uses `--nprocs` when calculating scheduler `total_steps`. Set it to the global `WORLD_SIZE` so the step count reflects data-parallel sharding, including multi-node runs.
|
||||
|
||||
For multi-GPU training, you **must** set these environment variables:
|
||||
Raw Slurm variables such as `SLURM_PROCID`, `SLURM_NTASKS`, and `SLURM_LOCALID` are not recognized automatically. Launch through `torchrun`, or map the scheduler's variables to `RANK`, `WORLD_SIZE`, `LOCAL_RANK`, `MASTER_ADDR`, and `MASTER_PORT` before starting AstrAI. The same requirement applies to launchers that expose only OpenMPI-specific variables.
|
||||
|
||||
## NCCL Troubleshooting
|
||||
|
||||
The following variables are troubleshooting options for hardware or network configurations where NCCL hangs or fails. They are not general requirements and can reduce performance by disabling peer-to-peer or GPUDirect RDMA paths:
|
||||
|
||||
```bash
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export NCCL_NET_GDR_LEVEL=0
|
||||
```
|
||||
|
||||
These are required on certain hardware configurations (see `AGENTS.md`). Without them, NCCL may hang or crash during collective operations. These are set in the training shell scripts (`train-seq.sh`, `train-sft.sh`, `train-dpo.sh`) but not in Python code — you must export them before launching.
|
||||
Apply them only after confirming the relevant NCCL transport is the source of the failure. AstrAI does not set them in Python.
|
||||
|
||||
## Checkpoint Saving
|
||||
|
||||
@@ -176,9 +187,6 @@ This ensures the LR schedule is correctly scaled regardless of the number of GPU
|
||||
|
||||
```bash
|
||||
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export NCCL_NET_GDR_LEVEL=0
|
||||
|
||||
python scripts/tools/train.py \
|
||||
--train_type=seq \
|
||||
--param_path ./params \
|
||||
@@ -240,7 +248,7 @@ python scripts/tools/train.py \
|
||||
|
||||
| Parameter | Default | Description |
|
||||
|-----------|---------|-------------|
|
||||
| `--nprocs` | 1 | Number of GPUs / processes |
|
||||
| `--nprocs` | 1 | Local process count for AstrAI's launcher; under `torchrun`, set it to global `WORLD_SIZE` for total-step calculation |
|
||||
| `--parallel_mode` | `fsdp` | `none`, `ddp`, or `fsdp` |
|
||||
| `--start_method` | `spawn` | Multiprocessing start method (`spawn`, `fork`, `forkserver`) |
|
||||
| `--backend` | `nccl` | Distributed backend (`nccl`, `gloo`) |
|
||||
@@ -248,8 +256,8 @@ python scripts/tools/train.py \
|
||||
| `--master_port` | `29500` | Master node port |
|
||||
| `--device_type` | `cuda` | Device type |
|
||||
|
||||
> `--tp_size` is parsed but **not yet wired** — tensor parallelism is future work. `ColumnParallelLinear` / `RowParallelLinear` exist in `astrai/parallel/module.py` but are not used by the model.
|
||||
> `--tp_size` is accepted by the CLI but discarded before configuration. Tensor parallelism is not implemented, and there is no tensor-parallel module or model integration.
|
||||
|
||||
Full parameter reference: [CLI Reference](params.md). Training loop and strategies: [Training Guide](training.md).
|
||||
|
||||
> Document Update Time: 2026-07-30
|
||||
> Document Update Time: 2026-08-02
|
||||
|
||||
+46
-12
@@ -2,6 +2,29 @@
|
||||
|
||||
AstrAI provides 7 evaluation scripts in `scripts/eval/` covering code generation, knowledge QA, perplexity, summarization, data quality, instruction following, and weight analysis.
|
||||
|
||||
## Contents
|
||||
|
||||
- [Prerequisites](#prerequisites)
|
||||
- [Overview](#overview)
|
||||
- [HumanEval](#humaneval-code-generation)
|
||||
- [MMLU](#mmlu-knowledge-qa)
|
||||
- [Perplexity](#perplexity-ppl)
|
||||
- [ROUGE](#rouge)
|
||||
- [IFD](#ifd-instruction-following-difficulty)
|
||||
- [IFEval](#ifeval-instruction-following)
|
||||
- [Weight Analysis](#weight-analysis)
|
||||
- [Tips](#tips)
|
||||
|
||||
## Prerequisites
|
||||
|
||||
HumanEval, MMLU, and IFEval import HuggingFace `datasets` to download their benchmark data. This package is not installed by AstrAI's base dependencies, so install it before running those scripts:
|
||||
|
||||
```bash
|
||||
pip install datasets
|
||||
```
|
||||
|
||||
The generation-based scripts require CUDA because they load the model on `cuda` with `bfloat16`. Direct-scoring and metric scripts support the devices shown below.
|
||||
|
||||
## Overview
|
||||
|
||||
| Script | Metric | Model Invocation | External Dataset |
|
||||
@@ -18,7 +41,15 @@ Two invocation patterns exist:
|
||||
- **Generation benchmarks** (HumanEval, IFEval): use `InferenceEngine` to generate responses, then score them.
|
||||
- **Scoring benchmarks** (MMLU, PPL, IFD): call `model()` directly under `torch.inference_mode()` for log-likelihood computation.
|
||||
|
||||
Common defaults: `--param_path` defaults to `./params`; dtype defaults to `bfloat16` on CUDA, `float32` on CPU.
|
||||
| Script | Device support |
|
||||
|--------|----------------|
|
||||
| HumanEval | CUDA for generation; `--test_only` can score existing completions without loading a model |
|
||||
| IFEval | CUDA only |
|
||||
| MMLU | CUDA or CPU via `--device`; auto-selects CUDA when available |
|
||||
| PPL | CUDA or CPU via `--device`; auto-selects CUDA when available |
|
||||
| IFD | CUDA or CPU via `--device`; auto-selects CUDA when available |
|
||||
| ROUGE | CPU-only metric computation; no model is loaded |
|
||||
| Weight analysis | CUDA by default; CPU supported via `--device cpu` |
|
||||
|
||||
---
|
||||
|
||||
@@ -30,7 +61,7 @@ Generates completions for 164 programming problems, executes them against hidden
|
||||
python scripts/eval/evaluate_humaneval.py \
|
||||
--param_path ./params \
|
||||
--num_samples 20 \
|
||||
--batch_size 32 \
|
||||
--batch_size 64 \
|
||||
--max_tokens 512 \
|
||||
--output results/humaneval.json
|
||||
```
|
||||
@@ -47,7 +78,8 @@ python scripts/eval/evaluate_humaneval.py \
|
||||
| `--temperature` | 0.8 | Sampling temperature |
|
||||
| `--top_p` | 0.95 | Nucleus sampling threshold |
|
||||
| `--top_k` | 50 | Top-k sampling |
|
||||
| `--batch_size` | 32 | Generation batch size |
|
||||
| `--batch_size` | 64 | Generation batch size |
|
||||
| `--max_seq_len` | 4096 | KV cache sequence length |
|
||||
| `--test_workers` | 8 | ProcessPoolExecutor workers for test execution |
|
||||
| `--test_timeout` | 3.0 | Per-subprocess timeout (seconds) |
|
||||
| `--problems` | None | Restrict to specific problem indices |
|
||||
@@ -66,7 +98,7 @@ python scripts/eval/evaluate_humaneval.py \
|
||||
python scripts/eval/evaluate_mmlu.py \
|
||||
--param_path ./params \
|
||||
--n_shot 5 \
|
||||
--subjects math_algebra history_us \
|
||||
--subjects abstract_algebra high_school_us_history \
|
||||
--output results/mmlu.json
|
||||
```
|
||||
|
||||
@@ -82,12 +114,13 @@ python scripts/eval/evaluate_mmlu.py \
|
||||
| `--device` | auto | Device (`cuda` / `cpu`) |
|
||||
| `--dtype` | auto | `bfloat16` on CUDA, `float32` on CPU |
|
||||
| `--seed` | 0 | Seed for option permutation (0 = enabled, -1 = disabled) |
|
||||
| `--batch_size` | 4 | Questions per batch; each question produces four choice rows |
|
||||
|
||||
**How it works**: For each question, builds a prompt with n-shot examples, then scores each choice (A/B/C/D) by computing the summed log-likelihood of the choice token given the context. The choice with the highest log-prob is the prediction.
|
||||
|
||||
**Output**: stdout prints per-subject accuracy and overall. With `--output`, writes per-subject `{accuracy, correct, total}` + `_overall` aggregate.
|
||||
|
||||
**Data**: Auto-downloads `cais/mmlu` from HuggingFace. Stored as per-subject CSVs in `<data_dir>/<split>/` and `<data_dir>/dev/` (for few-shot).
|
||||
**Data**: Auto-downloads `cais/mmlu` from HuggingFace. Stored as per-subject CSVs in `<data_dir>/<split>/` and `<data_dir>/dev/` (for few-shot). `--subjects` accepts canonical MMLU names such as `abstract_algebra`, `college_computer_science`, `high_school_us_history`, and `world_religions`.
|
||||
|
||||
---
|
||||
|
||||
@@ -100,7 +133,7 @@ python scripts/eval/evaluate_ppl.py \
|
||||
--param_path ./params \
|
||||
--input_path data.jsonl \
|
||||
--output_dir ppl_results/ \
|
||||
--batch_size 4 \
|
||||
--batch_size 64 \
|
||||
--max_length 2048
|
||||
```
|
||||
|
||||
@@ -110,7 +143,7 @@ python scripts/eval/evaluate_ppl.py \
|
||||
| `--input_path` | required | Input file, glob, or directory |
|
||||
| `--output_dir` | required | Output directory for `summary.json` + token JSONL |
|
||||
| `--text_key` | `text` | Key for the text field in input data |
|
||||
| `--batch_size` | 4 | Batch size |
|
||||
| `--batch_size` | 64 | Batch size |
|
||||
| `--max_length` | 2048 | Max sequence length (tokens) |
|
||||
| `--token_level` | False | Store per-token log_probs + token-type analysis |
|
||||
| `--max_samples` | None | Random subsample per file |
|
||||
@@ -119,7 +152,7 @@ python scripts/eval/evaluate_ppl.py \
|
||||
|
||||
**Input**: JSONL or JSON files. Each item must have a field named by `--text_key` (default `text`). If `--input_path` is a directory, recursively collects `*.jsonl` and `*.json`.
|
||||
|
||||
**Output**: `summary.json` with per-file stats (tokens, mean/median loss, perplexity, p50/p90/p95/p99). With `--token_level`, also writes per-token JSONL with token IDs and log-probs.
|
||||
**Output**: `summary.json` with per-file token count, mean loss, perplexity, and p50/p90/p95/p99 loss. Median loss is included only with `--token_level`; that mode also writes per-token JSONL with token IDs and log-probs.
|
||||
|
||||
---
|
||||
|
||||
@@ -210,7 +243,8 @@ python scripts/eval/evaluate_ifeval.py \
|
||||
| `--top_p` | 0.95 | Top-p sampling |
|
||||
| `--top_k` | 50 | Top-k sampling |
|
||||
| `--num_samples` | 1 | Samples per problem (best-of-n scoring) |
|
||||
| `--batch_size` | 1 | Inference batch size |
|
||||
| `--batch_size` | 64 | Inference batch size |
|
||||
| `--max_seq_len` | 4096 | KV cache sequence length |
|
||||
| `--limit` | None | Limit to first N problems (quick testing) |
|
||||
| `--dump_responses` | None | Path to dump raw responses as JSONL |
|
||||
|
||||
@@ -232,7 +266,7 @@ python scripts/eval/analyze_weights.py \
|
||||
|
||||
| Parameter | Default | Description |
|
||||
|-----------|---------|-------------|
|
||||
| `--ckpt_dir` | required | Checkpoint dir with `model.safetensors` + `config.json` |
|
||||
| `--ckpt_dir` | required | Checkpoint directory containing `model.safetensors` |
|
||||
| `--compare` | None | Additional checkpoint dirs to compare |
|
||||
| `--no_svd` | False | Skip SVD; show only weight stats (faster) |
|
||||
| `--output` | None | Save results as JSON |
|
||||
@@ -245,8 +279,8 @@ python scripts/eval/analyze_weights.py \
|
||||
## Tips
|
||||
|
||||
- **Quick test**: Use `--limit` (IFEval) or `--problems` (HumanEval) to run on a small subset first.
|
||||
- **Auto-download**: HumanEval, MMLU, and IFEval auto-download their datasets on first run. The other scripts expect user-provided data.
|
||||
- **Auto-download**: After installing `datasets`, HumanEval, MMLU, and IFEval auto-download their datasets on first run. The other scripts expect user-provided data.
|
||||
- **Output formats**: `--output` writes a single JSON for most scripts. PPL and IFD write an `--output_dir` containing `summary.json` plus per-file artifacts.
|
||||
- **CPU mode**: All scripts auto-detect CUDA. To force CPU, use `--device cpu --dtype float32`.
|
||||
- **CPU mode**: MMLU, PPL, and IFD support `--device cpu --dtype float32`; weight analysis supports `--device cpu`. HumanEval generation and IFEval are CUDA-only.
|
||||
|
||||
> Document Update Time: 2026-07-30
|
||||
|
||||
+52
-13
@@ -49,8 +49,7 @@ KVCache
|
||||
├── seq_lens [batch_size]
|
||||
├── out_cache_loc [batch, seq_len] — write indices for this forward
|
||||
├── max_len int — max(seq_lens), avoids GPU sync in decode
|
||||
├── page_table [batch, max_len] — precomputed gather indices for decode (None for prefill)
|
||||
└── decode_mask [batch, max_len] bool — precomputed position validity mask (None for single-batch decode)
|
||||
└── kv_indptr [batch + 1] int32 — prefix sum of seq_lens, precomputed once per step
|
||||
```
|
||||
|
||||
Attention layers do raw buffer indexing: `k_buffer[layer_id, out_cache_loc] = k` to write, `k_buffer[layer_id, indices]` to gather.
|
||||
@@ -87,7 +86,9 @@ Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/rotary_b
|
||||
- **CUDA kernel** (`rotary_emb.cu`): fused cos/sin lookup + rotation in a single kernel, used when the kernel is available, input is on CUDA, and `torch.is_grad_enabled()` is `False` (inference mode)
|
||||
- **Torch fallback**: complex multiply path (`torch.view_as_complex` → `torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd backward) or when the CUDA kernel is not available
|
||||
|
||||
`RotaryEmbedding` stores `cos_table`/`sin_table` as f32 buffers and returns a `(cos, sin)` tuple from `forward()`. Both attention backends share the same rotary dispatch — it is backend-agnostic.
|
||||
`RotaryEmbedding` stores a complex `freqs_cis` buffer and returns a tensor
|
||||
from `forward()`. Both attention backends share the same rotary dispatch — it
|
||||
is backend-agnostic.
|
||||
|
||||
## Continuous Batching
|
||||
|
||||
@@ -183,22 +184,58 @@ curl -X POST http://localhost:8000/v1/messages \
|
||||
-d '{"model":"astrai","system":"You are helpful.","messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
|
||||
```
|
||||
|
||||
Supports `stop_sequences` and streaming via `event: content_block_delta`.
|
||||
Supports `stop_sequences` and streaming via `event: content_block_delta`. Anthropic streams also end with the shared `data: [DONE]` sentinel after `event: message_stop`.
|
||||
|
||||
### GenerationRequest Parameters
|
||||
### Request Parameters
|
||||
|
||||
The HTTP protocols and direct engine API have distinct request models and defaults.
|
||||
|
||||
**OpenAI** (`ChatCompletionRequest`):
|
||||
|
||||
| Param | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `model` | str | `"astrai"` | Model name returned in responses |
|
||||
| `messages` | List[dict] | required | Chat messages (role, content) |
|
||||
| `top_k` | int | 50 | Top-k count |
|
||||
| `top_p` | float | 1.0 | Nucleus threshold |
|
||||
| `temperature` | float | 1.0 | Sampling temperature (> 0.0) |
|
||||
| `max_tokens` | Optional[int] | None | Max generation length |
|
||||
| `stream` | bool | False | Stream output |
|
||||
| `temperature` | Optional[float] | 1.0 | Sampling temperature (0.0-2.0) |
|
||||
| `top_p` | Optional[float] | 1.0 | Nucleus threshold (0.0-1.0) |
|
||||
| `top_k` | Optional[int] | 50 | Top-k count |
|
||||
| `max_tokens` | Optional[int] | 2048 | Max generation length |
|
||||
| `stream` | Optional[bool] | False | Stream output |
|
||||
| `stop` | Optional[Union[str, List[str]]] | None | Stop sequences |
|
||||
| `frequency_penalty` | float | 0.0 | Frequency penalty |
|
||||
| `tools` | Optional[List[dict]] | None | Tool definitions for function calling |
|
||||
| `tool_choice` | Optional[str] | None | Tool selection mode |
|
||||
| `n` | Optional[int] | 1 | Number of choices requested |
|
||||
| `presence_penalty` | Optional[float] | 0.0 | Presence penalty (-2.0 to 2.0) |
|
||||
| `frequency_penalty` | Optional[float] | 0.0 | Frequency penalty (-2.0 to 2.0) |
|
||||
| `logit_bias` | Optional[Dict[int, float]] | None | Per-token logit bias |
|
||||
| `user` | Optional[str] | None | End-user identifier |
|
||||
| `tools` | Optional[List[ToolDef]] | None | Tool definitions for function calling |
|
||||
| `tool_choice` | Optional[Union[str, Dict[str, Any]]] | `"auto"` | Tool selection mode or explicit tool choice |
|
||||
|
||||
**Anthropic** (`MessagesRequest`):
|
||||
|
||||
| Param | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `model` | str | `"astrai"` | Model name returned in responses |
|
||||
| `messages` | List[AnthropicMessage] | required | User/assistant messages |
|
||||
| `system` | Optional[str] | None | System prompt |
|
||||
| `max_tokens` | int | 1024 | Max generation length |
|
||||
| `temperature` | Optional[float] | 1.0 | Sampling temperature (0.0-2.0) |
|
||||
| `top_p` | Optional[float] | 1.0 | Nucleus threshold (0.0-1.0) |
|
||||
| `top_k` | Optional[int] | 50 | Top-k count |
|
||||
| `stream` | Optional[bool] | False | Stream output |
|
||||
| `stop_sequences` | Optional[List[str]] | None | Stop sequences |
|
||||
|
||||
**Engine** (`GenerationRequest`):
|
||||
|
||||
| Param | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `messages` | List[Dict[str, str]] | required | Messages to format before generation |
|
||||
| `top_k` | int | 50 | Top-k count; 0 disables filtering |
|
||||
| `top_p` | float | 1.0 | Nucleus threshold |
|
||||
| `temperature` | float | 1.0 | Sampling temperature; 0 enables greedy decoding |
|
||||
| `max_tokens` | Optional[int] | None | Max generation length |
|
||||
| `frequency_penalty` | float | 0.0 | Frequency penalty (-2.0 to 2.0) |
|
||||
| `rep_window` | int | 64 | Recent-token window used by the frequency penalty |
|
||||
| `stream` | bool | False | Stream output |
|
||||
|
||||
### SSE Streaming Format
|
||||
|
||||
@@ -240,6 +277,8 @@ data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":
|
||||
|
||||
event: message_stop
|
||||
data: {"type":"message_stop"}
|
||||
|
||||
data: [DONE]
|
||||
```
|
||||
|
||||
### Error Responses
|
||||
|
||||
+37
-30
@@ -13,9 +13,11 @@
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--config`, `-c` | YAML config file; explicit CLI options override YAML values | None |
|
||||
| `--train_type` | Training type (`seq`, `sft`, `dpo`, `grpo`, `online_grpo`, `online_dpo`) | required |
|
||||
| `--data_root_path` | Dataset root directory | required |
|
||||
| `--param_path` | Model parameters or checkpoint path | required |
|
||||
| `--resume` | Resume training from `--param_path` | False |
|
||||
| `--n_epoch` | Total training epochs | 1 |
|
||||
| `--batch_per_device` | Batch size per device | 1 |
|
||||
| `--grad_accum_steps` | Gradient accumulation steps between optimizer steps | 1 |
|
||||
@@ -26,7 +28,7 @@
|
||||
|-----------|-------------|---------|
|
||||
| `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 |
|
||||
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
|
||||
| `--max_grad_norm` | Maximum gradient norm for clipping (None disables) | 1.0 |
|
||||
| `--max_grad_norm` | Maximum gradient norm for clipping; the current CLI requires a positive number | 1.0 |
|
||||
|
||||
### Optimizer
|
||||
|
||||
@@ -36,9 +38,9 @@ non-matrix parameters through **AdamW** (`fused=True`).
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--optimizer` | Built-in optimizer (`muon_adamw`, `nora_nadamw`, `mano_adamw`) | `muon_adamw` |
|
||||
| `--weight_decay` | Weight decay (applied to Muon matrix params; non-matrix use 0) | 0.1 |
|
||||
| `--weight_decay` | Weight decay for optimizer parameter groups that are eligible for decay | 0.1 |
|
||||
| `--muon_momentum` | Muon momentum factor | 0.95 |
|
||||
| `--muon_nesterov` | Enable Nesterov momentum for Muon | True |
|
||||
| `--muon_nesterov`, `--no-muon_nesterov` | Enable or disable Nesterov momentum for Muon | enabled |
|
||||
| `--muon_ns_steps` | Newton-Schulz iteration steps for Muon | 5 |
|
||||
| `--muon_adjust_lr` | Muon LR adjustment strategy (`original`, `match_rms_adamw`) | `match_rms_adamw` |
|
||||
|
||||
@@ -56,15 +58,18 @@ under DTensor sharding and rejects layouts sharded along the last dimension.
|
||||
| `--nora_weight_decay` | Nora matrix weight decay | 0.0 |
|
||||
|
||||
`mano_adamw` routes internal `Linear.weight` matrices to **Mano** (manifold
|
||||
normalized optimizer) and the remaining parameters to **NAdamW**. Mano projects
|
||||
normalized optimizer) and the remaining parameters to **AdamW**. Mano projects
|
||||
the momentum onto the tangent space of the Oblique manifold and normalizes it,
|
||||
alternating the projection axis (row/column) each step — replacing Muon's
|
||||
Newton-Schulz iteration with a cheaper normalization.
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--mano_momentum` | Mano momentum factor | 0.95 |
|
||||
| `--mano_nesterov` | Enable Nesterov momentum for Mano | True |
|
||||
| `--mano_momentum` | Accepted by the CLI but currently ignored by optimizer construction | 0.95 |
|
||||
| `--mano_nesterov`, `--no-mano_nesterov` | Accepted by the CLI but currently ignored by optimizer construction | enabled |
|
||||
|
||||
The two Mano-specific flags are reserved for future wiring; do not rely on them
|
||||
to change optimizer behavior in the current release.
|
||||
|
||||
Optimizer identity and hyperparameters are saved in checkpoint metadata. Optimizer
|
||||
states are intentionally not interchangeable: resume older MuonAdamW checkpoints
|
||||
@@ -78,7 +83,7 @@ with `--optimizer=muon_adamw`.
|
||||
| `--stride` | Stride for sliding window over sequences | None |
|
||||
| `--random_seed` | Random seed for reproducibility | 3407 |
|
||||
| `--num_workers` | DataLoader worker processes | 4 |
|
||||
| `--no_pin_memory` | Disable pin_memory (enabled by default) | (flag) |
|
||||
| `--pin_memory`, `--no-pin_memory` | Enable or disable DataLoader pinned memory | enabled |
|
||||
|
||||
### Checkpoint & Resume
|
||||
|
||||
@@ -100,14 +105,20 @@ with `--optimizer=muon_adamw`.
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--log_dir` | Directory for metric logs | checkpoint/logs |
|
||||
| `--metrics` | Metrics to log (e.g. --metrics loss lr val_loss) | ["loss", "lr", "grad_norm"] |
|
||||
| `--metrics` | Repeatable metric option (for example, `--metrics loss --metrics lr --metrics val_loss`) | `loss`, `lr`, `grad_norm`, `grad_snr` |
|
||||
|
||||
### Gradient Checkpointing
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--gradient_checkpointing` | Enable activation checkpointing for DecoderBlock modules | False |
|
||||
| `--gradient_checkpointing`, `--no-gradient_checkpointing` | Enable or disable activation checkpointing for DecoderBlock modules | disabled |
|
||||
|
||||
### Miscellaneous
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--compile` | Enable `torch.compile` with mode `default`, `reduce-overhead`, or `max-autotune`; omit to disable | None |
|
||||
| `--dry-run` | Validate the merged configuration and print the training plan without training | False |
|
||||
|
||||
### Distributed Training
|
||||
|
||||
@@ -120,21 +131,25 @@ with `--optimizer=muon_adamw`.
|
||||
| `--backend` | Distributed training backend | nccl |
|
||||
| `--master_addr` | Master node address | localhost |
|
||||
| `--master_port` | Master node port | 29500 |
|
||||
| `--tp_size` | Reserved tensor-parallel size; accepted but currently ignored | None |
|
||||
|
||||
### Strategy-specific
|
||||
|
||||
| Parameter | Description | Default | Used by |
|
||||
|-----------|-------------|---------|---------|
|
||||
| `--dpo_beta` | DPO beta value | 0.1 | `dpo` |
|
||||
| `--dpo_beta` | DPO beta value | 0.1 | `dpo`, `online_dpo` |
|
||||
| `--label_smoothing` | Label smoothing for cross-entropy loss | 0.0 | `seq`, `sft` |
|
||||
| `--group_size` | GRPO group size | 4 | `grpo` |
|
||||
| `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo` |
|
||||
| `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 | `grpo` |
|
||||
| `--group_size` | GRPO/rollout group size | 4 | `grpo`, `online_grpo`, `online_dpo` |
|
||||
| `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo`, `online_grpo` |
|
||||
| `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 | `grpo`, `online_grpo` |
|
||||
| `--neftune_alpha` | NEFTune noise alpha (0=disabled, typical: 5.0) | 0.0 | `sft` |
|
||||
|
||||
### Online Rollout
|
||||
|
||||
These options apply to `online_grpo` and `online_dpo`. Online strategies require
|
||||
`online_grpo` and `online_dpo` are factory aliases for the existing `grpo` and
|
||||
`dpo` strategy classes; online behavior is enabled by rollout components rather
|
||||
than separate strategy subclasses. These options apply to the online aliases.
|
||||
Online strategies require
|
||||
a `BaseRewardModel` factory in `TrainConfig`; `train.py` does not currently
|
||||
provide a command-line option for configuring one.
|
||||
|
||||
@@ -151,7 +166,7 @@ provide a command-line option for configuring one.
|
||||
| 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.05 for cosine/SGDR, 0.0 for WSD) |
|
||||
| `--min_rate` | Minimum LR as fraction of base LR | None (all current schedulers use their effective default of 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 (80% of post-warmup steps) |
|
||||
@@ -204,14 +219,6 @@ python scripts/tools/server.py --param_path ./params --device cuda --dtype bfloa
|
||||
|
||||
See [Inference Guide](inference.md) for HTTP API documentation.
|
||||
|
||||
# Preprocess
|
||||
|
||||
```bash
|
||||
python scripts/tools/preprocess.py data/*.jsonl -o output/ -c config.json
|
||||
```
|
||||
|
||||
See [Preprocessing Guide](preprocessing.md) for config file format and examples.
|
||||
|
||||
## Generate (`generate.py`)
|
||||
|
||||
| Parameter | Type | Default | Description |
|
||||
@@ -221,13 +228,12 @@ See [Preprocessing Guide](preprocessing.md) for config file format and examples.
|
||||
| `--output_json_file` | str | required | Path to the output JSONL file |
|
||||
| `--question_key` | str | `question` | Key for the question in input JSON |
|
||||
| `--response_key` | str | `response` | Key for the response in output JSON |
|
||||
| `--temperature` | float | `0.60` | Sampling temperature |
|
||||
| `--top_k` | int | `30` | Top-k filtering |
|
||||
| `--temperature` | float | `0.8` | Sampling temperature |
|
||||
| `--top_k` | int | `50` | Top-k filtering |
|
||||
| `--top_p` | float | `0.95` | Nucleus sampling threshold |
|
||||
| `--batch_size` | int | `1` | Batch size for generation |
|
||||
| `--num_samples` | int | `1` | Responses per prompt |
|
||||
| `--max_tokens` | int | model config `max_position_embeddings` | Maximum tokens to generate |
|
||||
| `--cache_len` | int | `2048` | KV cache length |
|
||||
| `--max_seq_len` | int | `2048` | KV cache sequence length |
|
||||
| `--frequency_penalty` | float | `0.0` | Frequency penalty |
|
||||
| `--rep_window` | int | `64` | Window size for frequency penalty |
|
||||
|
||||
@@ -243,14 +249,15 @@ python scripts/tools/generate.py \
|
||||
|
||||
| Parameter | Type | Default | Description |
|
||||
|-----------|------|---------|-------------|
|
||||
| `input_files` | path(s) | required | Input JSONL file(s), supports glob (`data/*.jsonl`) |
|
||||
| `input_files` | path(s) | required | One or more existing `.jsonl` or `.json` paths. Wildcards work only when expanded by the invoking shell; the CLI does not expand globs itself. |
|
||||
| `--output_dir`, `-o` | path | required | Output directory for processed data |
|
||||
| `--config`, `-c` | path | required | Preprocessing pipeline config (JSON) |
|
||||
| `--tokenizer_path` | str | `params` | Path to tokenizer directory |
|
||||
| `--batch_size` | int | config value (`256` by default) | Override records processed per batch; must be at least 1 |
|
||||
|
||||
Usage:
|
||||
```bash
|
||||
python scripts/tools/preprocess.py data/*.jsonl -o output/ -c sft.json
|
||||
python scripts/tools/preprocess.py data/part-000.jsonl data/part-001.jsonl -o output/ -c sft.json
|
||||
```
|
||||
|
||||
See [Preprocessing Guide](preprocessing.md) for config file format and examples.
|
||||
|
||||
@@ -10,6 +10,7 @@ Declarative JSON-driven data preprocessing. `MaskBuilderFactory` supports three
|
||||
- [Configuration Reference](#configuration-reference) — all fields
|
||||
- [Mask Algorithm](#mask-algorithm)
|
||||
- [Output Layout](#output-layout)
|
||||
- [Training Compatibility](#training-compatibility)
|
||||
- [CLI](#cli)
|
||||
- [Python API](#python-api)
|
||||
|
||||
@@ -40,7 +41,7 @@ A single config file captures the entire pipeline, reusable and version-controll
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `field` | str | -- | JSONL key to read |
|
||||
| `action` | str | -- | `"train"` / `"mask"` / `"$role"` |
|
||||
| `action` | str | -- | `"train"` / `"mask"` / `"$role"` / `"value"`; `"value"` copies raw values without tokenization |
|
||||
| `template` | bool | `false` | Apply `chat_template` per message |
|
||||
| `add_special_tokens` | bool | `true` for first non-template section | Add special tokens during encode |
|
||||
|
||||
@@ -89,7 +90,7 @@ Config:
|
||||
}
|
||||
```
|
||||
|
||||
Output keys: `sequence` (int32), `loss_mask` (bool)
|
||||
Output keys: `sequence` (int32), `loss_mask` (bool), `position_ids` (int32)
|
||||
|
||||
### SFT Instruction
|
||||
|
||||
@@ -116,7 +117,7 @@ Config:
|
||||
}
|
||||
```
|
||||
|
||||
Output keys: `sequence`, `loss_mask`
|
||||
Output keys: `sequence`, `loss_mask`, `position_ids`
|
||||
|
||||
### Pretrain
|
||||
|
||||
@@ -142,7 +143,7 @@ Config:
|
||||
}
|
||||
```
|
||||
|
||||
Output keys: `sequence` (no `loss_mask` — all tokens trained)
|
||||
Output keys: `sequence`, `position_ids` (no `loss_mask` — all tokens trained)
|
||||
|
||||
### DPO
|
||||
|
||||
@@ -180,6 +181,11 @@ Config:
|
||||
|
||||
Output keys: `chosen`, `chosen_mask`, `rejected`, `rejected_mask`
|
||||
|
||||
The offline `Pipeline` can construct these keys, but its `.bin` output is not
|
||||
currently loadable for DPO training because the writer does not preserve
|
||||
per-record offsets. Train DPO directly from raw JSONL instead; see
|
||||
[Training Compatibility](#training-compatibility).
|
||||
|
||||
### GRPO
|
||||
|
||||
Input JSONL:
|
||||
@@ -228,6 +234,11 @@ Output keys: `prompts`, `prompts_mask`, `responses`, `masks`, `rewards` (float32
|
||||
- `mask_key: "masks"` — rename the auto-generated mask key (default: `responses_mask`)
|
||||
- `prompts_mask` is auto-generated (all masked) and unused by GRPOStrategy
|
||||
|
||||
The offline `Pipeline` flattens GRPO response groups for `.bin` output without
|
||||
preserving their boundaries, and there is no automatic raw-JSONL GRPO processor
|
||||
in `DatasetFactory`. See
|
||||
[Training Compatibility](#training-compatibility) for the supported routes.
|
||||
|
||||
---
|
||||
|
||||
## Configuration Reference
|
||||
@@ -257,7 +268,7 @@ When `sources` is set, `sections` is ignored.
|
||||
| `max_chars` | int | `2000000` | Skip text-mode items longer than this |
|
||||
| `max_items` | int or null | `null` | Stop after N documents |
|
||||
| `batch_size` | int | `256` | Records per tokenization batch |
|
||||
| `packing_strategy` | str | `"simple"` | Packing strategy: `"simple"`, `"bfd"`, `"bfd_split"` |
|
||||
| `packing_strategy` | str | `"simple"` | Packing is supported for single-output data with a `sequence` key: `"simple"`, `"bfd"`, or `"bfd_split"`. Multi-output DPO/GRPO data is not record-preserving packed output. |
|
||||
| `max_packed_len` | int | `8192` | Maximum length of a packed bin |
|
||||
| `truncation_mode` | str | `"keep_start"` | How to truncate sequences: `"keep_start"` or `"keep_end"` |
|
||||
|
||||
@@ -266,8 +277,8 @@ When `sources` is set, `sections` is ignored.
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `domain_key` | str or null | `null` | JSONL key for domain grouping |
|
||||
| `storage_format` | str | `"bin"` | `"bin"` (mmap). Reading also supports `"jsonl"` for on-the-fly tokenization |
|
||||
| `max_tokens_per_shard` | int | `100000000` | Flush threshold in cumulative tokens |
|
||||
| `storage_format` | str | `"bin"` | Pipeline output format. Only `"bin"` has a registered writer; `"jsonl"` is accepted by config validation but cannot be emitted by `Pipeline`. |
|
||||
| `max_tokens_per_shard` | int | `100000000` | Flush threshold counted from each record's primary flat sequence: `sequence` for single-output data, otherwise the first flat source output |
|
||||
| `dtype` | dict[str, str] | `{}` | Per-key tensor dtype override (e.g. `{"loss_mask": "bool"}`) |
|
||||
| `position_ids_mode` | str | `"doc_reset"` | How to compute position_ids: `"none"`, `"doc_reset"`, `"continuous"` |
|
||||
|
||||
@@ -304,11 +315,13 @@ output/
|
||||
meta.json
|
||||
sequence.bin
|
||||
loss_mask.bin
|
||||
position_ids.bin
|
||||
wiki/
|
||||
shard_0000/
|
||||
meta.json
|
||||
sequence.bin
|
||||
loss_mask.bin
|
||||
position_ids.bin
|
||||
```
|
||||
|
||||
### Multi-Shard (`bin`)
|
||||
@@ -322,13 +335,44 @@ output/
|
||||
meta.json
|
||||
sequence.bin
|
||||
loss_mask.bin
|
||||
position_ids.bin
|
||||
shard_0001/
|
||||
meta.json
|
||||
sequence.bin
|
||||
loss_mask.bin
|
||||
position_ids.bin
|
||||
```
|
||||
|
||||
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.
|
||||
`MmapStore` discovers binary shards recursively through their `meta.json` files.
|
||||
Each shard's metadata is a top-level object keyed by tensor name:
|
||||
|
||||
```json
|
||||
{
|
||||
"sequence": {"shape": [123456], "dtype": "int32"},
|
||||
"loss_mask": {"shape": [123456], "dtype": "bool"},
|
||||
"position_ids": {"shape": [123456], "dtype": "int32"}
|
||||
}
|
||||
```
|
||||
|
||||
An optional `offsets` array may appear for record-oriented binary data written
|
||||
through `save_bin(..., record_keys=...)`; the preprocessing `BinWriter` does not
|
||||
currently request those offsets.
|
||||
|
||||
---
|
||||
|
||||
## Training Compatibility
|
||||
|
||||
| Training type | Supported input route |
|
||||
|---------------|-----------------------|
|
||||
| `seq` | Offline preprocessed `.bin`, or raw `.jsonl` eagerly transformed by `JsonlStore` using `dataset_config.json` or the built-in `messages` config |
|
||||
| `sft` | Offline preprocessed `.bin`, or raw `.jsonl` through the same eager transform routes |
|
||||
| `dpo` | Raw `.jsonl` through the automatic lazy DPO processor selected by `DatasetFactory` when `tokenizer_path` is supplied, or a caller-provided record store |
|
||||
| `grpo` | A caller-provided, already-loaded `Store` with record-shaped `prompts`, `responses`, `masks`, and `rewards`; no automatic raw-JSONL processor is currently wired |
|
||||
|
||||
Offline DPO and GRPO preprocessing configs describe the intended token fields,
|
||||
but their `.bin` output is not currently loadable for training. DPO binary
|
||||
shards lack per-record offsets. GRPO response groups are flattened before the
|
||||
binary writer and their record/group boundaries are not preserved.
|
||||
|
||||
---
|
||||
|
||||
@@ -336,15 +380,20 @@ For `bin` format, `MmapStore` discovers all shards under the domain directory vi
|
||||
|
||||
```bash
|
||||
# SFT
|
||||
python scripts/tools/preprocess.py data/sft/*.jsonl -o output/sft/ -c configs/sft_chat.json
|
||||
python scripts/tools/preprocess.py data/sft/part-000.jsonl -o output/sft/ -c configs/sft_chat.json --batch_size 128
|
||||
|
||||
# DPO
|
||||
python scripts/tools/preprocess.py data/dpo/*.jsonl -o output/dpo/ -c configs/dpo.json --tokenizer_path params
|
||||
python scripts/tools/preprocess.py data/dpo/part-000.jsonl -o output/dpo/ -c configs/dpo.json --tokenizer_path params
|
||||
|
||||
# GRPO
|
||||
python scripts/tools/preprocess.py data/grpo/*.jsonl -o output/grpo/ -c configs/grpo.json
|
||||
python scripts/tools/preprocess.py data/grpo/part-000.jsonl -o output/grpo/ -c configs/grpo.json
|
||||
```
|
||||
|
||||
Inputs may be `.jsonl` files or `.json` files containing one object or a list of
|
||||
objects. Each positional path must exist. A wildcard such as `data/*.jsonl`
|
||||
works only when the invoking shell expands it before Click receives the
|
||||
arguments; otherwise pass the files explicitly.
|
||||
|
||||
---
|
||||
|
||||
## Python API
|
||||
|
||||
+18
-12
@@ -41,7 +41,10 @@ RoPE embeds position into Q/K vectors via complex rotation:
|
||||
|
||||
$$ q_i = R_i W_q x_i, \quad k_j = R_j W_k x_j, \quad q_i^T k_j = x_i^T W_q^T R_{i-j} W_k x_j $$
|
||||
|
||||
`RotaryEmbedding` pre-computes `cos_table` and `sin_table` (f32, `[max_len, dim/2]`). `forward()` returns a `(cos, sin)` tuple indexed by `position_ids`. `apply_rotary_emb` applies the rotation: during training it uses torch complex multiply (autograd-compatible); during inference it auto-dispatches to a fused CUDA kernel when available.
|
||||
`RotaryEmbedding` pre-computes a complex `freqs_cis` buffer. `forward()` returns
|
||||
a tensor indexed by `position_ids`. `apply_rotary_emb` applies the rotation:
|
||||
during training it uses torch complex multiply (autograd-compatible); during
|
||||
inference it auto-dispatches to a fused CUDA kernel when available.
|
||||
|
||||
## Training Loop
|
||||
|
||||
@@ -52,8 +55,8 @@ on_train_begin
|
||||
model.train()
|
||||
on_epoch_begin
|
||||
for batch in dataloader:
|
||||
on_batch_begin
|
||||
with executor.accumulate(model):
|
||||
on_batch_begin
|
||||
loss_output = strategy(batch)
|
||||
context.loss = loss_output["loss"].item()
|
||||
context.metrics = loss_output["metrics"]
|
||||
@@ -67,6 +70,7 @@ on_train_begin
|
||||
if executor.sync_gradients:
|
||||
on_optimizer_step
|
||||
optimizer.step()
|
||||
strategy.on_optimizer_step()
|
||||
optimizer.zero_grad()
|
||||
if scheduler:
|
||||
scheduler.step()
|
||||
@@ -78,16 +82,16 @@ on_train_end
|
||||
|
||||
| Hook | Fires | Default callback |
|
||||
|------|-------|-----------------|
|
||||
| `on_train_begin` | Before training starts | `GradientCheckpointingCallback` |
|
||||
| `on_train_begin` | Before training starts | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` |
|
||||
| `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` |
|
||||
| `on_batch_begin` | Every batch | — |
|
||||
| `on_optimizer_step` | Every accumulation window | `GradientClippingCallback`, `MetricCallback`, `ProgressBarCallback` |
|
||||
| `on_optimizer_step` | Every accumulation window | `MetricCallback`, `ProgressBarCallback`, `GradientClippingCallback` |
|
||||
| `on_batch_end` | Every batch | `CheckpointCallback` |
|
||||
| `on_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` |
|
||||
| `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` |
|
||||
| `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricCallback`, `GradientCheckpointingCallback` |
|
||||
| `on_train_end` | Training exits after `on_train_begin` completes (via `finally`) | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` |
|
||||
|
||||
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`).
|
||||
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm, rank-0), `gradient_clipping`. The gradient-clipping callback is always registered and always calls `executor.clip_grad_norm()` with the numeric `max_grad_norm` value.
|
||||
|
||||
Strategies return `{"loss": Tensor, "metrics": Dict[str, float]}` when called by the trainer. Built-in metrics include the task-specific loss and, for MoE models, `moe_aux_loss` plus `moe_aux_loss_weighted`. Direct `compute_loss(batch)` calls continue to return a single loss tensor.
|
||||
|
||||
@@ -98,7 +102,7 @@ Strategies return `{"loss": Tensor, "metrics": Dict[str, float]}` when called by
|
||||
Next-token cross-entropy with optional label smoothing:
|
||||
|
||||
$$
|
||||
L_{\text{PT}} = -\sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta)
|
||||
L_{\text{PT}} = -\frac{1}{T}\sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta)
|
||||
$$
|
||||
|
||||
Keys: `input_ids`, `target_ids`. Optional: `label_smoothing`.
|
||||
@@ -108,7 +112,7 @@ Keys: `input_ids`, `target_ids`. Optional: `label_smoothing`.
|
||||
Masked cross-entropy (`ignore_index=-100`) over response tokens:
|
||||
|
||||
$$
|
||||
L_{\text{SFT}} = -\sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta)
|
||||
L_{\text{SFT}} = -\frac{1}{L}\sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta)
|
||||
$$
|
||||
|
||||
Keys: `input_ids`, `target_ids`, `loss_mask`, `position_ids`. Optional: `label_smoothing`.
|
||||
@@ -168,9 +172,9 @@ model factory.
|
||||
|------|-------|-------------|
|
||||
| Cosine | `CosineScheduler` | Linear warmup → cosine decay to `min_rate` |
|
||||
| SGDR | `SGDRScheduler` | Cosine annealing with warm restarts (`t_mult=2`) |
|
||||
| WSD | `WSDScheduler` | Warmup-Stable-Decay with sqrt cooldown |
|
||||
| WSD | `WSDScheduler` | Warmup-Stable-Decay with quadratic decay |
|
||||
|
||||
Created by `SchedulerFactory.create(schedule_type, optimizer, **kwargs)`. Valid types: `"cosine"`, `"sgdr"`, `"wsd"`. Omit to use no scheduler.
|
||||
Created by `SchedulerFactory.create(schedule_type, optimizer, **kwargs)`. Valid types: `"cosine"`, `"sgdr"`, `"wsd"`. The training CLI always creates a scheduler and defaults `--schedule_type` to `"cosine"`.
|
||||
|
||||
## Gradient Checkpointing
|
||||
|
||||
@@ -188,10 +192,12 @@ Callback wraps each `DecoderBlock.forward` with `torch.utils.checkpoint.checkpoi
|
||||
|
||||
```
|
||||
Checkpoint(state_dict, epoch, consumed_samples, extra, meta, config)
|
||||
├── 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)
|
||||
├── save(save_dir) 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
|
||||
```
|
||||
|
||||
`Checkpoint.save()` writes whenever it is called. During training, `CheckpointCallback` uses the executor checkpoint context so only rank 0 receives a state dict and calls `save()`.
|
||||
|
||||
Optimizer/scheduler state persisted by default via `Checkpoint.extra`.
|
||||
Model config (`context.model_config`) saved into `config.json` during training via `CheckpointCallback`.
|
||||
|
||||
@@ -235,4 +241,4 @@ nohup python scripts/tools/train.py \
|
||||
|
||||
Full parameter reference at [params.md](params.md).
|
||||
|
||||
> Document Update Time: 2026-07-31
|
||||
> Document Update Time: 2026-08-02
|
||||
|
||||
Reference in New Issue
Block a user