Imported from Gabz4200/NonLinearRNNsCanBeParallel (
AGENTS.md). Install upstream withnpx skills add Gabz4200/NonLinearRNNsCanBeParallel. Copyright stays with the author.
AGENTS.md
Project overview
Research codebase for Nonlinear RNNs Can Be Parallel (arXiv:2603.03612): train nonlinear recurrent models in parallel chunks (scaffold + translator + target RNN) with parity against plain BPTT, evaluated on graph-connectivity tasks and causal language modeling.
- Stack: Python 3.11+ (pinned
3.13.14, upper bound<3.14), PyTorch 2.6, Lightning 2.3, Hydra/OmegaConf, Hugging Facetransformers/datasets, uv, ruff, pyrefly, pytest. - Layout:
src/nonlinearrnnscanbeparallel/(models, data, tasks, losses, training, integrations),configs/(Hydra),scripts/(train/eval/export),tests/,notebooks/(jupytextpy:percent),results/(run artifacts). - No monorepo, no CI workflows (
.github/absent). Pre-commit runs ruff + basic file hooks.
Human-facing docs live in README.md. This file is only what agents need to work here.
Setup
Requires uv. CPU and CUDA extras are mutually exclusive — pick one:
uv sync --extra cpu --extra dev # CPU boxes
uv sync --extra cu124 --extra dev # CUDA 12.4 boxes
Optional: --extra notebooks for jupyter/jupytext. Verify install with uv run pytest -q.
After any pyproject.toml edit: uv lock then uv lock --check.
Commands
Training (Hydra)
# smoke-test a config first (few batches) — always do this for new/changed configs
uv run python scripts/train.py --config-name config_wiki_tiny fast_dev_run=true
# default config: sorted_graph_connectivity, mlp_rnn, parallel
uv run python scripts/train.py
# parity check: both modes in one run
uv run python scripts/train.py --config-name config_wiki_tiny modes=[parallel,bptt]
# dotted overrides for anything else
uv run python scripts/train.py --config-name config_ustcon \
models=[mlp_rnn] modes=[parallel] data.num_samples=64 trainer.max_epochs=1
Other entry points:
uv run python scripts/eval.py— graph-connectivity test-split eval.uv run python scripts/export_hf.py --checkpoint <ckpt> --output-dir <dir>— Lightning → HF safetensors. Known broken on transformers 5.x:BaseModelOutputmust be imported fromtransformers.modeling_outputs, not the package root — fix before relying on export.uv run python scripts/param_count.py— param budget check forconfig_lmdefaults; exits 1 if target is outside 124M–350M (not a pure printer).notebooks/01_quickstart.py— jupytext notebook; keeppy:percentformat (see.jupytext.toml).
Each run writes results/<task>_<timestamp>/ (resolved config, plots, summary.json, RUN_REPORT.md) and Hydra logs under outputs/. Both are gitignored — never commit them.
Quality gates (run before delivery, all from repo root)
uv run ruff check --fix .
uv run ruff format .
uv run pyrefly check
uv run pytest
aislop scan .
- Focused iteration:
uv run pytest tests/test_parallel_wrapper.pyoruv run pytest -k "<name>". Full suite before delivery (~80s, 120 tests). aislop scan .uses the globalaislopbinary (not a project dependency). Fix errors and fixable warnings; don't edit.aislop/config without user consent.- Known baseline:
ruff checkcurrently reports ~34 pre-existing N806/N812/N803 findings (ML shape unpackingB, T, Dandimport torch.nn.functional as Fvs pep8-naming), andpyrefly checkreports ~14 pre-existing errors. Do not mass-rename unrelated code to silence them. Fix findings in code you touch; introduce no new ones.
Pre-commit
uv run pre-commit install
uv run pre-commit run --all-files
Hooks: trailing whitespace, EOF, yaml check, large files, ruff --fix, ruff-format.
- Known:
pre-commit run --all-filescurrently fails — pinned hook ruffv0.6.9differs from project ruff0.16.7(e.g. flagsUP038, reformats some asserts). Prefer the Quality gates commands above for delivery; only chase pre-commit failures you introduce.
Testing
- Location/convention:
tests/test_*.py;pythonpath = ["src"]; default addopts-v(see[tool.pytest.ini_options]inpyproject.toml). - Naming:
test_when_<scenario>_then_<outcome>. Annotate test functions-> None. One Act per test.pytest.raisesfor expected exceptions;@pytest.mark.parametrizewith descriptiveids. - Mock only external boundaries; prefer real model code. No network, no sleeps, deterministic seeds.
- New model code requires (mirror existing patterns in
tests/test_models.py):- reference
forwardon CPU (shape + backward), - full-sequence
forwardvs token-by-tokenstepconsistency, - if touching
parallel_wrapper.pyor trainer paths: a parallel-vs-BPTT parity case (seetests/test_parallel_lm_parity.py).
- reference
Code style
- Ruff: line length 100, target py311, rules
E,F,W,I,N,UP,B,C4,SIM. Formatter:uv run ruff format .. - Types: pyrefly (
pyrefly.toml); coverssrc/,tests/,scripts/,benchmarks/. Excludes.venv/,outputs/,notebooks/,integrations/. - Layout: package under
src/; scripts stay thin and import from the package (they alsosys.path.insertsrc/so plainpython scripts/...works). - Imports: stdlib → third-party → local; ruff isort rules enforce order.
- Fail fast: no silent fallbacks or broad
exceptin library code; let programmer errors raise. - Comments only for non-obvious intent/invariants (why a reshape, gradient-isolation reason, paper reference). No narrative comments, no commented-out code, no decorative separators.
- Tests and public docstrings use normal full-clarity prose.
Architecture notes
Model protocol
Every cell implements BaseRNNModel (models/base.py):
forward(x, state) -> (logits, new_state)— full sequence[B, T, D]step(x_t, state) -> (logits, new_state)— single token[B, D], O(1) decodeinit_state(batch_size, device) -> RNNStateList
State travels in RNNState / RNNStateList dataclasses (hidden + optional extra), not bare tensors. Keep mutation/cloning/detach explicit.
Parallel training path
models/parallel_wrapper.py (ParallelRNNTrainer) decomposes each layer into scaffold (minGRU/minLSTM, parallel-scannable) → translator (MLP or rKAN) → target RNN (in-chunk BPTT). Gradients stay isolated across the boundary path. If parallel.chunk_size >= sequence length, run falls back to single-chunk sequential and warns — that is a config bug, not a feature; lower chunk_size (e.g. 64 for 256-token blocks).
Adding a model
- Implement the cell protocol in
src/nonlinearrnnscanbeparallel/models/. - Decorate with
@register_model("<name>")(models/registry.py). - Add a preset under
configs/model/<name>_<size>.yamland wire it into aconfigs/config_*.yamlmodels:list. - Cover it in
tests/test_models.py(forward/backward +stepparity). import nonlinearrnnscanbeparallel.modelshas a registry side effect — needed in any entry point that callsget_model.
Config system
Hydra configs compose model / data / trainer / parallel defaults (configs/config.yaml and configs/config_*.yaml). Top-level knobs: task, models (list), modes (parallel | bptt or both), fast_dev_run. Parallel knobs in configs/parallel.yaml (chunk_size, scaffold_*, translator_*). Prefer CLI dotted overrides over editing shared configs for one-off runs.
Tasks and metrics
Tasks: sorted_graph_connectivity, ustcon, long_sequence, openthoughts_lm, wikipedia_lm. LM metrics: train/loss, val/loss, val/nll, val/ppl, val/acc (val/nll aliases cross-entropy). Central claim under test is parity between parallel and bptt — small sub-percent loss gaps are normal (independent weight init/shuffle unless seeded); third-decimal-and-beyond gaps need investigation.
Build and deployment
No packaged deployment. Research loop only:
- Install/sync:
uv sync --extra cpu --extra dev(orcu124). - Export for HF consumers:
scripts/export_hf.py→safetensors+ config underintegrations/transformers/. benchmarks/holds throughput benchmarks (currently empty scaffolding).
PR / commit guidelines
- No required CI; local gates above are the bar. Run them before asking for review.
- Commit message: conventional style (
feat:,fix:,test:,chore:, …), imperative mood, subject ≤50 chars when practical. - Never commit:
.env, secrets,.venv,outputs/,lightning_logs/, generatedresults/*/,*.ckpt,.aislop/,.opencode/,.omp(see.gitignore). - Keep diffs scoped; don't reformat or rename unrelated code while fixing a bug.
Gotchas
--extra cpuand--extra cu124conflict — pick one oruv syncfails.uv run aislopis wrong (not a dependency); callaislop scan ..- Scripts must run as
uv run python scripts/train.py ...(Hydraconfig_path=../configsdepends on that entry point). - Parallel run checkpoint weights:
backbone.target.*(parallel) vsbackbone.*(BPTT). - Generation from trained models is experimental only (collocations/loops at current scale) — load target weights into
NanoRNNand decode withstep(). - Notebook edits: keep
notebooks/*.pyas jupytext percent scripts; don't hand-edit orphaned.ipynb. - Don't run
pre-commit run --all-filesexpecting a clean pass on untouched files (version skew with project ruff — see Pre-commit). It rewrites EOF newlines across many files; revert unrelated churn.