Hacktoberfest 2026:メンテナが10月に向けて印を付けた、オープンで初心者向けの issue。 Hacktoberfest の issue を見る

[Bug Report] get_caching_hooks slices the wrong axis when remove_batch_dim=True is combined with pos_slice

クローズ
#1,856 コメント 1 件 リアクション 0 件 担当者 1 名 GitHub で見る

メンテナーはふだん 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, dev at 1012730f; same code on main.
  • 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_native model.

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

環境構築

Codespaces で開く

このプロジェクトの開発コンテナを、あなたの GitHub アカウントでブラウザ上に起動します。

  • Dockerfile・Docker Compose ファイルなし
  • プルリクエストのテンプレートあり
  • コントリビューションガイドなし

はじめの一歩

  1. issue を最後まで読み、次にプロジェクトのコントリビューションガイドを読みます。
  2. 着手することを issue にコメントします — 二人が同じ作業をするのを防げます。
  3. リポジトリをフォークし、ブランチを切って変更します。
  4. issue 番号を参照したプルリクエストを送ります。

TransformerLensOrg/TransformerLens のほかの issue

TransformerLensOrg/TransformerLens の issue をすべて見る

似ている issue

Python の issue をもっと見る

新しい issue をメールで受け取る

初心者向けの GitHub issue を短くまとめたダイジェスト。