git clone https://github.com/jon123boss/LR-AttnRes
cd LR-AttnResInstall required dependencies via pip:
pip install flash-attn --no-build-isolation
pip install tiktoken
pip install huggingface-hub
pip install datasets
pip install lm_eval
pip install hf_transfer
pip install wandb # Optional, for experiment trackingProvision the end-to-end qualified H100 stack with a complete managed Python
3.12 runtime (including Python.h):
python scripts/bootstrap_fast_attnres.py
source .venv/bin/activateThe bootstrap selects a driver-compatible profile. Driver 580 or newer uses Python 3.12, PyTorch 2.10.0+cu130, Triton 3.6.0, and FlashAttention 2.8.3's official CUDA-13/Torch-2.10 wheel. CUDA 12.x drivers use PyTorch 2.9.0+cu126, Triton 3.5.0, and the official CUDA-12/Torch-2.9 FlashAttention wheel. Both profiles install the same SHA-256-pinned Fast-AttnRes 2.0.1 wheel. Using the wheels avoids local FlashAttention builds; managed Python supplies the header Triton's runtime launcher needs.
train.py defaults to --attnres_backend auto. Standard AttnRes (R=D) and
single-head static sliced/output-tail LR-AttnRes (1<=R<=D) resolve to Fast.
Startup then requires the exact package/provenance and proves every
multi-source read is Fast before arming the model; every read checks that
contract again. A package, CUDA/BF16, shape, or semantic mismatch is a hard
error, never a silent legacy fallback. The first embedding-only read is an
identity and does not invoke any routing operator.
Fast-AttnRes v2.0.1 has no source-prior API, and the qualified compiled sliced path uses neutral scale 1.0. Therefore Block runs must disable count priors and sliced runs must disable LR logit scaling:
# Standard Full, R=D: automatically Fast.
python train.py --use_attnres --attnres_type full
# Sliced Block, R=64<D: automatically Fast or fails before training.
python train.py --use_lrid --lrid_key_from_output_tail \
--lrid_rank 64 --attnres_type block \
--no-attnres_block_count_prior --no-lrid_logit_scaleProjected keys, dynamic queries, multi-head LRID, and arbitrary source priors
are different equations and resolve to legacy under auto. Explicit
--attnres_backend fast makes them fail with the exact incompatibility;
explicit legacy remains available only for controlled reference comparisons.
Rank zero prints [Fast-AttnRes] ... only after the strict route is armed.
The reproducible H100 comparison (legacy, Fast, and pre-norm-only) is:
uvx --from modal==1.5.4 modal run modal_fast_attnres_benchmark.pyIt fixes public main's existing SDPA fallback across all arms and uses
torch.compile(fullgraph=True, dynamic=False) without max-autotune.
For reproducible downstream evaluation, use the tested, pinned environment instead of the unpinned commands above:
pip install -r requirements-eval.txtrequirements-eval-lock-macos-py38.txt records the exact transitive package
set used for the verified macOS/Python 3.8 CPU smoke run. Use a separately
generated platform lock for CUDA/Linux because PyTorch's platform dependencies
differ.
This repo now defaults to GPT-4-tokenized Ultra-FineWeb-en shards:
- tokenizer:
tiktoken.encoding_for_model("gpt-4")/cl100k_base - vocab size:
100277 - document separator token:
100257 - shard dtype:
uint32
Prepare 20B total tokens once and optionally upload them to your Hugging Face dataset repo:
python prepare_ultrafineweb.py \
--hf-repo-id <your-hf-username>/Ultra-FineWeb-en-20B-gpt4 \
--uploadAfter the shards are uploaded, download them for training:
python prepdata.py --repo-id <your-hf-username>/Ultra-FineWeb-en-20B-gpt4Single-process training still works with:
python train.pyFor DDP, launch with torchrun:
torchrun --standalone --nproc_per_node=8 train.pytorchrun --standalone --nproc_per_node=2 train.py --torch-max-autotune --full_runFor 8-GPU DDP with PyTorch compile max-autotune:
torchrun --standalone --nproc_per_node=8 train.py --torch-max-autotuneBy default, DDP preserves the configured global batch size by dividing
grad_accum_steps across ranks, so the default 8 accumulation steps becomes
1 local accumulation step on 8 GPUs. Use --no-ddp_preserve_global_batch if
you want global batch size to scale with WORLD_SIZE.
Enable PyTorch compile max-autotune with:
python train.py --torch-max-autotuneMax-autotune writes TorchInductor/Triton autotune caches. By default this uses PyTorch/Triton's normal cache locations. To force a specific large/persistent disk:
torchrun --standalone --nproc_per_node=8 train.py \
--torch-max-autotune \
--torch_compile_cache_dir /workspace/LR-AttnRes/out/torchinductor_cacheFor a full automated run that prompts for Hugging Face sign-in/repo setup at
startup, trains, saves the final checkpoint, uploads it to Hugging Face as
final_model.pt, and then runs evaluation:
torchrun --standalone --nproc_per_node=8 train.py \
--torch-max-autotune \
--full_run \
--full_run_hf_repo_id <your-hf-username>/<model-repo>If --full_run_hf_repo_id is omitted, full_run prompts for it at startup.
Evaluation results from the automatic eval are saved to
out/full_run_eval_step:<step>.txt.
To resume after an interrupted run, leave --ckpt_file_name empty to pick the
highest-numbered ckpt_step:<step>.pt in out_dir:
torchrun --standalone --nproc_per_node=2 train.py --init_from resume --ckpt_file_name ""run_eval.py can load checkpoints produced by DDP training because checkpoints
save the unwrapped model state on rank 0. A normal eval run is single-process:
The downstream suite uses lm-evaluation-harness task definitions with a
zero-shot override. It is not the official OLMES protocol (which uses its own
curated prompts, scoring variants, and aggregation), so label reported numbers
as generic lm-eval zero-shot results rather than OLMES results.
python run_eval.py --ckpts out/ckpt_step:1000.ptEvery run_eval.py invocation atomically saves a human-readable text report,
a complete structured JSON companion, and a status JSON. The default report
name includes a UTC timestamp so a failed run cannot be confused with an older
result:
python run_eval.py --ckpts out/ckpt_step:1000.pt --results-file out/eval_results.txtThe default suite is strict: every checkpoint and task must exist and every
task must return nonempty metrics. --allow-skipped is an explicit opt-in for
partial dataset-script runs. --limit 1 --tasks piqa is useful for a clearly
labelled one-example smoke test; it must not be reported as a full benchmark.
The default suite contains eight tasks and omits WinoGrande and MMLU. They can
still be requested explicitly with --tasks winogrande or --tasks mmlu.
The declared primary metrics are length-normalized accuracy for ARC-Challenge,
ARC-Easy, HellaSwag, OpenBookQA and PIQA, and raw accuracy for BoolQ,
CommonsenseQA and Social-IQA. WinoGrande also uses raw accuracy when requested.
When requested explicitly, MMLU's overall group aggregate is saved and
displayed separately from its subject-level task metrics.
For multi-GPU validation loss, launch with torchrun:
torchrun --standalone --nproc_per_node=8 run_eval.py --validation-onlyValidation loss is sharded across ranks and reduced exactly. Downstream lm-eval
tasks run on rank 0 only; when both validation and downstream tasks are enabled,
run_eval.py tears down the process group after validation before rank 0 starts
the long task pass. Eval compile is opt-in:
python run_eval.py --ckpts out/ckpt_step:1000.pt --torch-max-autotuneNew training checkpoints are written atomically and include every DDP rank's RNG state plus the exact number of consumed local training batches. Resuming restores model, optimizer, scheduler, random state, and the next shuffled batch. Older checkpoints still load, but cannot reproduce the exact uninterrupted trajectory because they did not store RNG or in-epoch dataloader state. Trajectory-exact resume also requires the recorded ordered shard manifest, objective, optimizer, clipping and schedule settings to match. The shard manifest checks path, file identity, size and nanosecond modification time; the loader rejects mismatches instead of silently resuming from different data. Training preserves the historical loss path and computes cross-entropy directly from bfloat16 logits. Validation casts logits to float32 for stable reporting. New checkpoints also guard the PyTorch/CUDA/device/criterion runtime and every Attention-Residual kernel-selection environment variable. Bitwise trajectory claims still require the same deterministic hardware and software environment.
Checkpoint loading preserves an explicitly saved attnres_block_count_prior.
If it is missing from model_args, the loader uses the saved training config's
value, or False when both omit it, matching checkpoints trained before the
count prior existed. Evaluation, training resume, analysis, inspection and
HF import/resave share this migration; new-training defaults are unchanged.
Imported/resaved models record the resolved value explicitly. Copies already
resaved by older code with an injected True must be checked against their
original checkpoint metadata, since they look like intentionally prior-trained
models. Rerun affected checkpoint-based evaluations after updating; original
in-training validation is unaffected by this loading bug. A continuation trained
after an incorrect load must restart from its last unaffected checkpoint to
recover the intended recipe.
LR AttnRes can be enabled as a block Attention Residuals variant:
python train.py --use_lrid --attnres_type block --lrid_rank 64--use_lrid automatically enables use_attnres. LR AttnRes uses the same learned,
input-independent depth queries as normal Attention Residuals, but routes over
low-rank input-dependent source keys. An optional ablation can also emit
input-dependent depth queries with --lrid_input_dependent_query, changing LR
output projections from d + k to d + 2k; this uses a gated hybrid query
static_query + gate * dynamic_query. Depth routing can be split into multiple
heads with --lrid_num_heads; lrid_rank remains the total low-rank width.
Use --lrid_static_embedding_key to make the embedding source key a learned,
input-independent LR key instead of projecting it from token embeddings.
Use --lrid_add_static_embedding_key or --lrid_add_static_source_key to add
a learned static key to the computed embedding key or computed non-embedding
source keys.
Use --lrid_key_from_value to project LR keys from the source value or block
summary instead of fusing them into each output projection. This is unshared by
default, keeping a separate value-key projector per LR output module;
--lrid_key_from_value_shared uses one shared source-key projection. Use
--lrid_query_from_value to do the same for dynamic queries, with
--lrid_query_from_value_shared for the shared variant. Outside key/query
projections use stateless rms_norm(source_value) by default; disable it with
--no-lrid_key_value_norm.
Logit scaling defaults to 1 / sqrt(lrid_rank / lrid_num_heads);
disable it with --no-lrid_logit_scale or set it explicitly with --lrid_logit_scale.
Attention Residual key normalization, query normalization, and query initialization
are configurable via --attnres_key_norm, --attn_res_query_norm, and
--attn_res_query_init. Block summaries are averaged by default with
--attnres_block_average; this divides by the sublayer count, and
--attnres_block_average_mode sqrt divides by the square root of the count.
For explicit experiments, --attnres_block_alpha sets the block value exponent
in sum(v_i) / count^alpha, and --attnres_block_beta sets the logit prior
beta * log(count). The default legacy values preserve the old controls:
count averaging is alpha=1, sqrt averaging is alpha=0.5, and
--attnres_block_count_prior gives beta=1 unless disabled. Alpha and beta can
also be learned with --attnres_block_alpha_learned and
--attnres_block_beta_learned, scoped by {shared,per_residual,per_block}.
Use --attnres_block_split_sublayers to keep separate attention-output and
MLP-output summaries inside each compressed block. This preserves a coarse
source-type routing axis while still using far fewer sources than full AttnRes.
For a learnable alternative, --attnres_block_learned_scale gives each live
partial/completed block source its own scalar, initialized with
--attnres_block_learned_scale_init {count,sqrt,one}. To remove block value
scale directly, --attnres_block_value_norm applies stateless RMSNorm to each
block source value instead of scalar scaling.
See LR_ATTNRES.md for the full design note, parameter cost, stability rationale, and experiment matrix.