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

[Bug Report] `remove_batch_dim=True` with batch `size > 1`: the three caching paths disagree

オープン
#1,858 コメント 3 件 リアクション 0 件 担当者 1 名 GitHub で見る

メンテナーはふだん 1 日以内に返信

@Mudassiruddin7 がすでに取り組んでいます。

2026年10月6日 から。

評価

この issue はまだ評価されていません。

説明

bug complexity-simple TransformerBridge

Describe the bug

The docstrings say remove_batch_dim "only makes sense with batch_size=1 inputs", but only one of the three caching paths enforces that; the other two silently do different wrong things on batch > 1:

  • run_with_cache(..., remove_batch_dim=True) (ActivationCache path): raises AssertionError: Cannot remove batch dimension from cache with batch size 2. This is loud and arguably the right behavior.
  • run_with_cache(..., remove_batch_dim=True, return_cache_object=False): silently ignores remove_batch_dim, activations keep their batch dim.
  • get_caching_hooks(remove_batch_dim=True) / add_caching_hooks: save_hook does stored[0], silently discarding every example after the first. This is bad, as it quietly abandons data.

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, (2, 6))
name = "blocks.0.hook_out"

with torch.no_grad():
    try:
        bridge.run_with_cache(tokens, names_filter=name, remove_batch_dim=True)
    except AssertionError as e:
        print("ActivationCache:", e)
    _, plain = bridge.run_with_cache(
        tokens, names_filter=name, remove_batch_dim=True, return_cache_object=False
    )
    print("plain dict:", tuple(plain[name].shape))

cache, fwd_hooks, _ = bridge.get_caching_hooks(names_filter=name, remove_batch_dim=True)
with torch.no_grad(), bridge.hooks(fwd_hooks=fwd_hooks):
    bridge.forward(tokens)
print("get_caching_hooks:", tuple(cache[name].shape))

Output on dev (1012730f):

ActivationCache: Cannot remove batch dimension from cache with batch size 2
plain dict: (2, 6, 32)
get_caching_hooks: (6, 32)

Expected: one behavior across all three – the ActivationCache assert is the obvious candidate (an assert tensor.size(0) == 1 in save_hook and in the return_cache_object=False squeeze). EDIT: We have made some slight adjustments to this expected behavior, see discussion below

System Info

  • Source checkout installed with uv sync, dev at 1012730f.
  • macOS (Apple Silicon), CPU only; Python 3.12.12; torch 2.11.0; transformers 5.13.0.
  • No model downloads needed: reproduced with a random-init TransformerBridge.boot_native model.

Additional context

Found while verifying #1856. The silent [0] is inherited from the legacy HookedRootModule.get_caching_hooks, which has the same behavior. A fix should probably touch both.

Checklist
  • I have checked that there is no similar issue in the repo (required)
主要言語
Python
スター
3.9k
フォーク
708
平均マージ
1日 18時間
マージ済み PR(30日)
65

環境構築

Codespaces で開く

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

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

はじめの一歩

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

TransformerLensOrg/TransformerLens のほかの issue

TransformerLensOrg/TransformerLens の issue をすべて見る

似ている issue

Python の issue をもっと見る

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

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