Token Averaging
Collection
A comprehensive research framework for analyzing whether averaging adjacent tokens in a LLM can reduce the compute compared to standard model. • 23 items • Updated
YAML Metadata Warning:empty or missing yaml metadata in repo card
Check out the documentation for more information.
Checkpoint dump from the token averaging research project.
avg_50m_k8_learnable_posresultsloss_log.csvcheckpoints/final.ptcheckpoints/step_00350000.ptcheckpoints/step_00400000.ptcheckpoints/step_00450000.ptimport torch
from huggingface_hub import hf_hub_download
path = hf_hub_download('FAIRC/token-averaging-avg_50m_k8_learnable_pos', 'checkpoints/final.pt')
state = torch.load(path, map_location='cpu', weights_only=False)
model.load_state_dict(state['model']) # your OLMAveraged / OLMTransformerBody
print(state['step'], state['tokens_seen'], state['cumulative_flops'])
These are not Hugging Face transformers weights. Rebuild the
architecture from config.json → model_config (or from
experiments/chinchilla/model_configs.py in the source repo) and load
the raw state_dict.
{
"d_model": 512,
"n_heads": 8,
"n_layers": 8,
"context_len": 1024,
"averaging_k": 8,
"tie_embeddings": true,
"method_name": "learnable_pos_k8",
"lr": 0.0002,
"warmup_steps": 2000,
"target_tokens": 8144000000,
"n_params_approx": 50897408
}