Hacktoberfest 2026: những issue maintainer đã đánh dấu cho tháng Mười, đang mở và phù hợp người mới. Xem issue Hacktoberfest

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

Đã đóng
#1,858 3 bình luận 0 reaction 1 người được giao Xem trên GitHub

Maintainer thường phản hồi trong vòng 1 ngày

@Mudassiruddin7 đang làm issue này rồi.

Từ ngày 6/10/2026.

Đánh giá

Issue này chưa được đánh giá.

Mô tả

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)
Ngôn ngữ chính
Python
Star
3.9k
Fork
708
Merge trung bình
1 ngày 17 giờ
Pull request đã merge (30 ngày)
70

Chuẩn bị môi trường

Mở trong Codespaces

Khởi chạy dev container của dự án ngay trên trình duyệt, bằng tài khoản GitHub của bạn.

  • Không có Dockerfile hay tệp Docker Compose
  • Có mẫu pull request
  • Không có hướng dẫn đóng góp

Bắt đầu từ đâu

  1. Đọc hết issue, rồi đọc hướng dẫn đóng góp của dự án.
  2. Bình luận trên issue rằng bạn sẽ nhận — tránh hai người làm cùng một việc.
  3. Fork repository và làm thay đổi trên một nhánh.
  4. Mở pull request có tham chiếu số hiệu của issue.

Issue khác của TransformerLensOrg/TransformerLens

Tất cả issue của TransformerLensOrg/TransformerLens

Issue tương tự

Thêm issue về Python

Nhận issue mới trong hộp thư của bạn

Bản tóm tắt ngắn những issue GitHub phù hợp với người mới.