RoboTTT
Implementation of RoboTTT proposed by Yunfan Jiang et al. of Stanford and Nvidia.
Install
pip install robo-ttt
Usage
Basic Usage
Using MemoryKeyValueBind and TTTWrapper standalone with a 2-layer MLP memory:
import torch
from torch import nn
from robo_ttt import MemoryKeyValueBind, TTTWrapper
memory_network = nn.Sequential(
nn.Linear(512, 1024),
nn.GELU(),
nn.Linear(1024, 512)
)
memory = MemoryKeyValueBind(512, memory_network)
ttt_wrapper = TTTWrapper(512, memory = memory)
# attended action chunks over time (batch = 2, time = 5, seq_len = 4, dim = 512)
attended_action_chunks = torch.randn(2, 5, 4, 512)
out, next_fast_weights, _ = ttt_wrapper(attended_action_chunks)
assert out.shape == attended_action_chunks.shape
You can also drop in Fast Weight Product Key Memory (Tianyu Zhao & Llion Jones) directly via fwPKMWrapper:
import torch
from robo_ttt import fwPKMWrapper, TTTWrapper
from fast_weight_product_key_memory import fwPKM
pkm = fwPKM(
dim = 512,
num_memories = 256 * 256,
dim_queries_keys = 512,
dim_values = 512
)
memory = fwPKMWrapper(pkm)
ttt_wrapper = TTTWrapper(512, memory = memory)
attended_action_chunks = torch.randn(2, 5, 4, 512)
out, next_fast_weights, _ = ttt_wrapper(attended_action_chunks)
Full Policy Wrapper with MimicVideo
Wrapping a policy model (e.g. MimicVideo) with RoboTTT:
import torch
from torch import nn
from mimic_video import MimicVideo
from robo_ttt import RoboTTT, MemoryKeyValueBind, TTTWrapper
memory_network = nn.Sequential(
nn.Linear(512, 1024),
nn.GELU(),
nn.Linear(1024, 512)
)
memory = MemoryKeyValueBind(512, memory_network)
ttt_wrapper = TTTWrapper(512, memory = memory)
policy = MimicVideo(
dim = 512,
dim_video_hidden = 512,
depth = 2,
dim_head = 64,
heads = 8,
dim_action = 4,
dim_joint_state = 4
)
model = RoboTTT(
policy,
ttt_wrapper = ttt_wrapper,
ttt_module_paths = ('to_action_tokens',),
batch_time_arg = 'video_hiddens',
expand_time_args = ('prompt_token_ids',),
times_arg = 'time'
)
# inputs for sequence of t = 3 timesteps (batch = 2, time = 3)
video_hiddens = torch.rand(2, 3, 5, 512)
joint_state = torch.randn(2, 3, 4)
actions = torch.randn(2, 3, 32, 4)
prompt_token_ids = torch.tensor([[10, 20, 30, -1], [15, 25, -1, -1]])
# forward training pass with loss masking
loss_mask = torch.tensor([[True, False, True], [False, True, True]])
loss = model(
prompt_token_ids = prompt_token_ids,
video_hiddens = video_hiddens,
actions = actions,
joint_state = joint_state,
loss_mask = loss_mask
)
loss.backward()
# sampling / rollout one timestep at a time (auto_unsqueeze_time defaults to True)
init_video_hiddens = torch.randn(2, 5, 512)
init_joint_state = torch.randn(2, 4)
actions_t1, fast_weights1 = model.sample(
prompt_token_ids = prompt_token_ids,
video_hiddens = init_video_hiddens,
joint_state = init_joint_state,
steps = 4,
batch_size = 2,
auto_unsqueeze_time = True,
return_fast_weights = True
)
Citations
@article{jiang2026robottt0,
title = {RoboTTT: Context Scaling for Robot Policies},
author = {Yunfan Jiang and Yevgen Chebotar and Ruijie Zheng and Fengyuan Hu and Yunhao Ge and Jimmy Wu and Tianyuan Dai and Scott Reed and Li Fei-Fei and Yuke Zhu and Linxi "Jim" Fan},
year = {2026},
journal = {arXiv preprint arXiv: 2607.15275}
}
@misc{zhao2026fastweightproductkeymemory,
title = {Fast-weight Product Key Memory},
author = {Tianyu Zhao and Llion Jones},
year = {2026},
eprint = {2601.00671},
archivePrefix = {arXiv},
primaryClass = {cs.CL},
url = {https://arxiv.org/abs/2601.00671},
}
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
robo_ttt-0.1.5.tar.gz
(11.4 kB
view details)
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
robo_ttt-0.1.5-py3-none-any.whl
(10.8 kB
view details)
File details
Details for the file robo_ttt-0.1.5.tar.gz.
File metadata
- Download URL: robo_ttt-0.1.5.tar.gz
- Upload date:
- Size: 11.4 kB
- Tags: Source
- Uploaded using Trusted Publishing? No
- Uploaded via:
uv/0.8.17
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
d24384150b878b2a46a4fb3795e267aef21b360c6a33d8cfc6432a1d0851e8ee
|
|
| MD5 |
531ab40bcf0f1a731b2cdb40bc5806d0
|
|
| BLAKE2b-256 |
43e195afa247a3d381a9d577294e0b51799ec804541fa5ec6d1255b5d91ff106
|
File details
Details for the file robo_ttt-0.1.5-py3-none-any.whl.
File metadata
- Download URL: robo_ttt-0.1.5-py3-none-any.whl
- Upload date:
- Size: 10.8 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via:
uv/0.8.17
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
20fdd6bc819cfcfab6525c8378d882f9af6df3bc40afbdaeaa2614be5df48057
|
|
| MD5 |
79ff3502dd10d7cfaab662d587c9d919
|
|
| BLAKE2b-256 |
fab666eb027e2196fd8f5f726652d7cc61681d12e17539229fc4a67f387deb48
|