[Bug Report] get_caching_hooks slices the wrong axis when remove_batch_dim=True is combined with pos_slice
メンテナーはふだん 1 日以内に返信
@JoeyTan21 がすでに取り組んでいます。
2026年10月6日 から。
評価
この issue はまだ評価されていません。
説明
Describe the bug
TransformerBridge.get_caching_hooks(remove_batch_dim=True, pos_slice=...) (and add_caching_hooks, which delegates to it) returns activations sliced on the wrong axis. save_hook drops the batch dimension first and then applies pos_slice on _pos_slice_dim(name), which is the positive axis 1 for every activation except attention maps. Once the batch dim is gone, axis 1 is d_model for the residual stream and n_heads for head-split tensors, so e.g. blocks.0.hook_out is cached as [pos, 1] instead of [1, d_model]. Attention patterns are unaffected because their axis is -2.
run_with_cache with the same arguments slices first and removes the batch dim afterwards, so the two caching APIs silently disagree for the same inputs.
Code example
import torch
from transformer_lens.config import TransformerBridgeConfig
from transformer_lens.model_bridge import TransformerBridge
cfg = TransformerBridgeConfig(
d_model=32, d_head=16, n_heads=2, n_layers=1, n_ctx=8, d_vocab=16,
d_mlp=64, act_fn="gelu", normalization_type="LN", seed=0,
)
bridge = TransformerBridge.boot_native(cfg)
tokens = torch.randint(0, cfg.d_vocab, (1, 6))
names = ["blocks.0.hook_out", "blocks.0.attn.o.hook_in", "blocks.0.attn.hook_pattern"]
with torch.no_grad():
_, ref = bridge.run_with_cache(tokens, names_filter=names, remove_batch_dim=True, pos_slice=-1)
cache, fwd_hooks, _ = bridge.get_caching_hooks(names_filter=names, remove_batch_dim=True, pos_slice=-1)
with torch.no_grad(), bridge.hooks(fwd_hooks=fwd_hooks):
bridge.forward(tokens)
for name in names:
print(name, tuple(ref[name].shape), tuple(cache[name].shape))
Output on dev (1012730f):
blocks.0.hook_out (1, 32) (6, 1)
blocks.0.attn.o.hook_in (1, 2, 16) (6, 1, 16)
blocks.0.attn.hook_pattern (2, 1, 6) (2, 1, 6)
Expected: the get_caching_hooks shapes (and values) match run_with_cache, i.e. (1, 32), (1, 2, 16), (2, 1, 6).
Without remove_batch_dim=True both APIs agree, so the bug is specifically the order of the two operations in get_caching_hooks.save_hook (transformer_lens/model_bridge/bridge_core.py). Swapping them — slice on the batched tensor, then drop the batch dim, as run_with_cache._store already does — fixes the reproduced cases; I have a patch with a unit test ready and can open a PR against dev.
System Info
- Source checkout installed with
uv sync,devat1012730f; same code onmain. - macOS (Apple Silicon), CPU only; Python 3.12.13; torch 2.11.0; transformers 5.x.
- No model downloads needed: reproduced with a random-init
TransformerBridge.boot_nativemodel.
Additional context
The legacy HookedRootModule.get_caching_hooks used negative position dims (-2 / -3) and so was immune to this ordering. Searched issues/PRs for pos_slice, remove_batch_dim, get_caching_hooks; #574 / #578 (2024) concern the HookedTransformer-era pos_slice=None crash, not this.
Checklist
- I have checked that there is no similar issue in the repo (required)
Disclosure: found, reproduced and the draft fix tested locally with the help of an AI coding agent (Claude Code); reviewed before filing.
- 主要言語
- Python
- スター
- 3.9k
- フォーク
- 708
- 平均マージ
- 1日 17時間
- マージ済み PR(30日)
- 70
環境構築
このプロジェクトの開発コンテナを、あなたの GitHub アカウントでブラウザ上に起動します。
- Dockerfile・Docker Compose ファイルなし
- プルリクエストのテンプレートあり
- コントリビューションガイドなし
はじめの一歩
- issue を最後まで読み、次にプロジェクトのコントリビューションガイドを読みます。
- 着手することを issue にコメントします — 二人が同じ作業をするのを防げます。
- リポジトリをフォークし、ブランチを切って変更します。
- issue 番号を参照したプルリクエストを送ります。
TransformerLensOrg/TransformerLens のほかの issue
-
[Proposal] RoBERTa masked-LM adapter for TransformerBridge対応中かも @Canonik が今日担当しました。 オープンcomplexity-moderate new-architecture TransformerBridge
TransformerLensOrg/TransformerLens#1870 · コメント 1 件 · 担当者 1 名 ·
メンテナーはふだん 1 日以内に返信
-
[Proposal] Backward Lens: support gated MLP gate/ up/ down gradient factors対応中かも @janmenjayap が 8 日前に担当しました。 オープンcomplexity-moderate enhancement TransformerBridge
TransformerLensOrg/TransformerLens#1832 · 担当者 1 名 ·
メンテナーはふだん 1 日以内に返信
-
[Proposal] Sparse probing: optional groups argument so rows from one prompt can't straddle the split対応中かも @lorenzozanee が 13 日前に担当しました。 オープンcomplexity-simple enhancement help wanted TransformerBridge
難易度 4/5 3〜5日 初心者へのやさしさ 25/100
TransformerLensOrg/TransformerLens#1813 ·
メンテナーはふだん 1 日以内に返信
-
[Bug Report] _BLOCK_LIST_ATTRS hardcoded name list silently drops Raven's blocks from composition-score / head-label analysis対応中かも @LightWork666 が 17 日前に担当しました。 オープンbug complexity-moderate TransformerBridge
TransformerLensOrg/TransformerLens#1791 · コメント 2 件 · 担当者 1 名 ·
メンテナーはふだん 1 日以内に返信
-
[Proposal] SVD Circuits: singular-vector decomposition of a head's QK/ OV into causally-validated subfunctions対応中かも @janmenjayap が 29 日前に担当しました。 オープンcomplexity-high enhancement TransformerBridge
TransformerLensOrg/TransformerLens#1767 · コメント 3 件 · 担当者 1 名 ·
メンテナーはふだん 1 日以内に返信
TransformerLensOrg/TransformerLens の issue をすべて見る
似ている issue
-
難易度 1/5 1時間未満 初心者へのやさしさ 85/100
MystenLabs/MemWal#1163 · コメント 2 件 ·
メンテナーはふだん 1 日以内に返信
-
infertopics leaves new nodes without a topic when untopiced neighbours outnumber topiced ones対応中かも @moneebullah25 が今日担当しました。 オープン
難易度 2/5 1〜3時間 初心者へのやさしさ 72/100
メンテナーはふだん 1 日以内に返信
-
難易度 2/5 1〜3時間 初心者へのやさしさ 70/100
FinanceFlash/unvibecode#218 ·
メンテナーはふだん 1 日以内に返信
-
難易度 2/5 1〜3時間 初心者へのやさしさ 75/100
メンテナーはふだん 1 日以内に返信
-
難易度 2/5 1〜3時間 初心者へのやさしさ 70/100
NVIDIA/earth2studio#1241 ·
メンテナーはふだん 3 日以内に返信