Skip to content

Latest commit

 

History

105 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

Download Repository

git clone https://github.com/jon123boss/LR-AttnRes
cd LR-AttnRes

Prerequisites

Install 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 tracking

Provision 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/activate

The 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_scale

Projected 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.py

It 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.txt

requirements-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.

Data Preparation

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 \
  --upload

After the shards are uploaded, download them for training:

python prepdata.py --repo-id <your-hf-username>/Ultra-FineWeb-en-20B-gpt4

Training

Single-process training still works with:

python train.py

For DDP, launch with torchrun:

torchrun --standalone --nproc_per_node=8 train.py
torchrun --standalone --nproc_per_node=2 train.py --torch-max-autotune --full_run

For 8-GPU DDP with PyTorch compile max-autotune:

torchrun --standalone --nproc_per_node=8 train.py --torch-max-autotune

By 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-autotune

Max-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_cache

For 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 ""

Evaluation

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.pt

Every 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.txt

The 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-only

Validation 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-autotune

New 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

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.

About

No description, website, or topics provided.

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages