TaH2 🌐 Project · 📑 Paper · 🤗 HuggingFace
TaH2 improves test-time scaling by allocating extra latent iterations to tokens that benefit from deeper computation. It jointly trains the backbone and an iteration decider with lookahead depth supervision, using online labels that indicate whether another iteration improves prediction. On challenging AIME benchmarks, TaH2 improves the accuracy-compute slope by 53% over the non-looped baseline and raises peak accuracy by about 3.4 points at matched test-time compute.
@article{you2026tah2,
title={Improving Test-Time Scaling with Adaptive Looped Transformers},
author={You, Yichen and Fu, Tianyu and Feng, Aosong and Lv, Xingtai and Ning, Xuefei and Ding, Ning and Wang, Yu},
journal={arXiv preprint arXiv:2609.35748},
year={2026},
}TaH 🌐 Project · 📑 Paper · 🤗 HuggingFace
Think-at-Hard (TaH) improves LLM reasoning by running extra latent iterations only on hard tokens instead of all tokens. A lightweight decider and duo-causal attention enable targeted refinement while keeping full parallelism. TaH outperforms fixed two-iteration baselines by 8–11% while skipping 94% of second iterations, and also beats strong single-iteration Qwen3 models by 4–5%.
@article{fu2025tah,
title={Think-at-Hard: Selective Latent Iterations to Improve Reasoning Language Models},
author={Tianyu Fu and Yichen You and Zekai Chen and Guohao Dai and Huazhong Yang and Yu Wang},
journal={arXiv preprint arXiv:2511.08577},
year={2025},
}-
[2026/10] We released the TaH2 code, models, and training data.
-
[2026/09] We introduced TaH2 in Improving Test-Time Scaling with Adaptive Looped Transformers.
-
[2025/11] We released the TaH-plus-1.7B checkpoint. The model is finetuned from Qwen3-1.7B-Base using 100K samples from the OpenR1 dataset, capable of QA, math, and coding.
-
[2025/11] Our paper was featured as the #2 Paper of the Day on Huggingface Daily Papers
TaH and TaH2 require different dependency versions. Activate their separate environments before running the corresponding scripts.
This repository includes training recipes, evaluation tools, and an inference engine adapted from mini-SGLang, integrated under tah2/minisgl/.
Use Linux, Python 3.12, and CUDA GPUs. TaH2 uses tah2/, script/tah2/, and bash/. From the repository root:
python3.12 -m venv .venv-tah2
source .venv-tah2/bin/activate
pip install -e '.[tah2]'Download a checkpoint:
hf download nics-efc/TaH2-1.7B-max2 --local-dir models/TaH2-1.7B-max2Generate with the bundled inference engine:
from transformers import AutoTokenizer
from tah2.minisgl.core import SamplingParams
from tah2.minisgl.llm import LLM
model_path = "models/TaH2-1.7B-max2"
tokenizer = AutoTokenizer.from_pretrained(model_path)
prompt = tokenizer.apply_chat_template(
[{"role": "user", "content": "What is 1 + 1?"}],
tokenize=False, add_generation_prompt=True,
)
llm = LLM(model_path=model_path)
try:
result = llm.generate([prompt], SamplingParams(max_tokens=1024, temperature=0.6))
print(result[0]["text"])
finally:
llm.shutdown()| Checkpoint | Model |
|---|---|
| TaH2-1.7B-max2 | Adaptive iteration, maximum depth 2 |
| TaH2-1.7B-Standard | Single-iteration baseline; also loads with Transformers |
Start a server, then run evaluation in another terminal:
MODEL_PATH=models/TaH2-1.7B-max2 bash bash/launch_server.sh \
--tah_iter_threshold 0.5MODEL_PATH=models/TaH2-1.7B-max2 bash bash/eval_online.sh \
--datasets math500 amc23 olympiadbench aime25 aime26The server defaults to GPU 0 and http://127.0.0.1:30080. Set GPU, SERVER_HOST, PORT, and DISTRIBUTED_PORT to change its placement; set BASE_URL for the evaluation client. The server reads the iteration limit from the checkpoint; --tah_max_iter overrides it. Set the client's --tah_iter_threshold to the server threshold when labeling evaluation outputs. The engine README also describes its Python API and OpenAI-compatible endpoints.
Step 0: Prepare Data and Base Model
Download the prepared data from nics-efc/TaH2-amteam-tool:
hf download nics-efc/TaH2-amteam-tool --repo-type dataset --local-dir dataThe math, code, science, and tool-calling mixture provides train/ and eval/ splits for each model.
| Student | Directory | Train samples | Train tokens | Eval samples |
|---|---|---|---|---|
| Qwen3-1.7B | data/1.7b/ |
273,195 | 1,099,413,406 | 955 |
| Qwen3-4B | data/4b/ |
638,817 | 2,586,857,725 | 1,000 |
| Qwen3-8B | data/8b/ |
1,282,124 | 5,192,783,145 | 1,000 |
Load with datasets.load_from_disk("data/1.7b/train"); mask=1 marks assistant tokens for training. Recipes load Qwen/Qwen3-{1.7B,4B,8B}-Base automatically, or use a local path in model.name.
Data Sources and Regeneration
| Original dataset | Files used |
|---|---|
| a-m-team/AM-Qwen3-Distilled | math.jsonl, code.jsonl, science.jsonl |
| nvidia/Nemotron-Agentic-v1 | data/tool_calling.jsonl |
For 1.7B, Qwen3-8B regenerates assistant responses; tool results use reference replay or Qwen3-32B simulation. Key code: generation and tokenization, tool rollout. Input prompts are available at 1.7b-prompts/.
Serve Qwen3-8B at port 30080 and Qwen3-32B at port 30081, then run:
python script/tah2/data/regenerate.py generate --kind am \
--prompts data/1.7b-prompts/am/train.jsonl \
--urls http://127.0.0.1:30080 --output data/regenerated/am_train.jsonl
python script/tah2/data/regenerate.py generate --kind tool_calling \
--prompts data/1.7b-prompts/tool_calling/train.jsonl \
--urls http://127.0.0.1:30080 --sim-urls http://127.0.0.1:30081 \
--output data/regenerated/tool_train.jsonl
python script/tah2/data/regenerate.py build \
--inputs data/regenerated/am_train.jsonl data/regenerated/tool_train.jsonl \
--output data/regenerated/trainPrepared data: 1.7b/, 4b/, 8b/. The 4B/8B mixtures retain the original teacher responses: 8B uses the full pool; 4B uses a source-stratified subset with the same eval split.
To regenerate both splits, build eval first, then pass --exclude data/regenerated/eval when building train to remove duplicate eval samples.
| Recipe directory | Standard | Fixed loop-2 | TaH2 |
|---|---|---|---|
qwen3_1.7/ |
sft_base.yaml |
sft_fixed.yaml |
sft_tah.yaml, sft_tah_max4.yaml, sft_tah_max8.yaml |
qwen3_4b/ |
sft_base.yaml |
— | sft_tah.yaml |
qwen3_8b/ |
sft_base.yaml |
— | sft_tah.yaml |
TaH2 jointly optimizes the backbone, input updater, and decider. Posterior labels are generated online from next-token cross-entropy improvements; separate offline token labeling is unnecessary. The main recipes use DUO attention, Triton kernels, and stop_prob_mix. The fixed loop-2 recipe uses even_mix without decider supervision.
The 1.7B TaH2 recipes use global batch size 128, three epochs, learning rate 4e-5, and a 16,384-token packing budget.
For one node with eight GPUs, run:
NPROC=8 TP=1 CONFIG=script/tah2/recipes/qwen3_1.7/sft_tah.yaml bash bash/sft_tah.shChange CONFIG to select a recipe and adjust TP to fit the actual GPU memory usage. Update data.dp in the recipe accordingly (data.dp = NPROC / TP for one node).
Suggested starting points for one node with 8 H200 GPUs, BF16, and gradient checkpointing:
| Model | max_iter |
max_length |
Suggested TP |
GPU (memory per GPU) |
|---|---|---|---|---|
| Qwen3-1.7B | 1 | 16K | 1 | H200 (141GB) |
| Qwen3-1.7B | 2 | 16K | 1 | H200 (141GB) |
| Qwen3-1.7B | 4 | 16K | 1 | H200 (141GB) |
| Qwen3-1.7B | 8 | 16K | 2 | H200 (141GB) |
| Qwen3-4B | 1 | 16K | 2 | H200 (141GB) |
| Qwen3-4B | 2 | 16K | 2 | H200 (141GB) |
| Qwen3-8B | 1 | 16K | 2 | H200 (141GB) |
| Qwen3-8B | 2 | 16K | 2 | H200 (141GB) |
Checkpoints include model weights, tokenizer, and the recurrent components when enabled. Recipes use save_only_model: true; set it to false before training to save optimizer state for continuation.
Use Python 3.10 and activate a separate environment for tah/ and script/tah/:
python3.10 -m venv .venv-tah
source .venv-tah/bin/activate
pip install -e '.[tah]'For training and evaluation, install additional dependencies:
pip install -e '.[tah,training,evaluation]'For code generation evaluation, install evalplus
Note if you
git pulland the top-level package layout changes (e.g.__init__.pyis added or removed), re-runpip install -e '.[tah]'— the editable install caches the layout insite-packages/__editable___tah_*_finder.pyand stale state will silently droptah/__init__.py's re-exports.
python script/tah/playground/inference_example.py # quick demo (~1 min)
python script/tah/playground/inference_example.py --max-new-tokens 16384 # full reasoning chainThis script demonstrates TaH's selective latent iteration mechanism, with color-coded output showing the iteration count for each token.
python script/tah/evaluation/eval.py \
--eval_config ./script/tah/recipes/qwen3_1.7/eval_tah.yaml \
--model_path nics-efc/TaH-plus-1.7B \
--dataset_name gsm8k \
--backend tah \
--job_nums 8 \
--tp_size_per_job 1Key parameters:
--eval_config: Path to evaluation config file--model_path: Path to the model--dataset_name: Dataset name (supports gsm8k, math500, aime24, etc. Detailed configs can be found intah/evaluate/eval_configs/dataset_configs.json)--backend: Inference backend (tahfor TaH)--job_nums: Number of parallel jobs (one job pinstp_size_per_jobGPUs)--tp_size_per_job: Tensor parallel size per job--data_range N/--data_range start end: subset slice — handy for smoke tests--data_ids gsm8k_0,gsm8k_5: run only specific problem ids
The default recipe targets 8 GPUs (--job_nums 8). To sanity-check the pipeline on
one GPU in a couple of minutes, slice the dataset and shrink max_new_tokens:
# clone the recipe and shrink generation length
sed 's/max_new_tokens: 4096/max_new_tokens: 512/' \
script/tah/recipes/qwen3_1.7/eval_tah.yaml > /tmp/eval_tah_smoke.yaml
CUDA_VISIBLE_DEVICES=0 python script/tah/evaluation/eval.py \
--eval_config /tmp/eval_tah_smoke.yaml \
--model_path nics-efc/TaH-plus-1.7B \
--dataset_name gsm8k --backend tah \
--job_nums 1 --tp_size_per_job 1 \
--data_range 5 \
--output_dir /tmp/tah_eval_smokeThe TaH backend is a token-by-token Python loop intended for research; for serving
throughput, use --backend sglang or the dedicated minisgl-tah server.
The same script/tah/evaluation/eval.py accepts --backend hf (vanilla
AutoModelForCausalLM.generate — useful for non-TaH baselines) or
--backend sglang (sgl Engine for high-throughput serving). All three
backends share the same job-sharded driver under
tah/evaluate/jobs.py:allocate_gpus_and_run_jobs.
Training a TaH model consists of three stages:
1. Prepare training data
Use a reference model to generate hard token labels for the training and validation data:
# download the default subset of OpenR1-Math-220k
python script/tah/preparation/download.py
# filter and split
python script/tah/preparation/filter_split.py
# label the hard tokens
python script/tah/preparation/label.py \
--num_gpu 8 \
--dataset_path ./data/initial_data/openr1-math/train.jsonl \
--test_model_list Qwen/Qwen3-1.7B \
--output_path ./data/processed_data/openr1-math/1_7/train \
--max_input_length 10000
python script/tah/preparation/label.py \
--num_gpu 8 \
--dataset_path ./data/initial_data/openr1-math/eval.jsonl \
--test_model_list Qwen/Qwen3-1.7B \
--output_path ./data/processed_data/openr1-math/1_7/eval \
--max_input_length 10000 \2. (Optional) Prepare pruned model
For the TaH version, prune one layer from the base model to match the parameter count of the standard baseline (skip this step for TaH+ version):
python script/tah/preparation/prune.py \
--model Qwen/Qwen3-1.7B-Base \
--dataset ./data/processed_data/openr1-math/1_7/eval \
--output ./model/qwen3_1.7_base_pruned \
--num_prune 1The first stage uses fixed iteration labels for training:
python -m accelerate.commands.launch \
--config_file ./script/tah/recipes/accelerate_configs/zero2.yaml \
--num_processes 8 \
./script/tah/train/SFT_TaH.py \
--config ./script/tah/recipes/qwen3_1.7/sft_tah_step1.yamlKey configurations in Step1 (sft_tah_step1.yaml):
max_iter: 2— maximum number of iterations.iter_decider: "IterLabelDecider"— continue iff the per-token oracleiter_count_labels(derived frommismatch) say so. Used to teach the LoRA adapter on tokens marked "hard" by the labeller.adapter: "lora"— only LoRA is supported in tah-release.train_loss: "NextTokenPredLoss"— standard causal-LM cross-entropy.
Single-implementation hooks (input/output updaters, iter labels, adapter) are inlined into the wrapper — only iter_decider and train_loss are config-selectable.
The second stage trains the iteration decider:
python -m accelerate.commands.launch \
--config_file ./script/tah/recipes/accelerate_configs/zero2.yaml \
--num_processes 8 \
./script/tah/train/SFT_TaH.py \
--config ./script/tah/recipes/qwen3_1.7/sft_tah_step2.yamlKey configurations in Step2 (sft_tah_step2.yaml):
tah_model_path: Load the model trained in Step1iter_decider: "MLPIterDecider": Use MLP decider to automatically determine iterationstrain_loss: "IterDeciderLoss": Iteration decider loss functionfreeze_component: [model.simple_base_model]: Freeze model backbone
After two-stage training, the model can automatically decide when to perform latent reasoning iterations.
TaH/
├── tah/ # TaH model, LoRA training, and evaluation
├── tah2/
│ ├── model/ # recurrent model, decider, posterior labels, losses
│ ├── kernels/ # Triton recurrent attention
│ ├── train/ # FSDP2/TP training and checkpoint saving
│ ├── evaluate/ # inference backends and benchmark grading
│ ├── minisgl/ # bundled mini-SGLang engine and native kernels
│ └── utils/ # data preparation and serialization
├── script/
│ ├── tah/ # TaH preparation, training, evaluation, and recipes
│ └── tah2/ # TaH2 data, training, evaluation, and recipes
├── bash/ # TaH2 training, evaluation, and server launchers
└── pyproject.toml # separate tah/tah2 dependency selections
TaH and TaH2 are released under Apache-2.0. The bundled mini-SGLang engine retains its MIT license and the bundled NCCL header retains its NVIDIA license.
Explore more efficient LLM projects from us:
|
R2R
Token-level routing for reasoning LLMs |
C2C
Communicate through KV-Cache between LLMs |
FrF
Efficient video token reduction for LVLMs |
MoA
Mixture of sparse attention for LLMs |