wooyeolbaek/attention-map-diffusers

attention map tools for huggingface/diffusers

何であるか

attention‑map‑diffusers は、Hugging Face Diffusers ライブラリ上で実行される拡散モデルの内部のアテンションテンソルをキャプチャおよび可視化できる Python パッケージです。画像生成パイプライン(例:Stable Diffusion、FLUX)および動画生成パイプライン(例:Wan、CogVideoX)の両方で動作します。推論中にクロスアテンションおよびセルフアテンション行列を記録することで、トークンや画像パッチが互いにどのように影響し合うかを示すヒートマップを生成します。たとえば、プロンプトのどの部分が生成画像の特定の領域を駆動しているか、または動画フレームが以前のフレームにどのように注目しているかを可視化できます。

なぜ重要なのか

拡散モデルはブラックボックス型の生成器です。モデルがどこに注目しているかを理解することは、研究者がプロンプトをデバッグし、モデルの挙動を研究し、より良い条件付け技術を開発する上で重要です。このライブラリは、元のモデルコードを変更せずに、その検査を簡単に行えるようにします。

コア機能(README に記載された内容)

機能 詳細
関係を意識したキャプチャ text→imageimage→imagevideo→textvideo→video など、記録するアテンション関係を指定できます。
多数のモデルに対応 数十のチェックポイントで事前にテスト済み:FLUX.2‑Klein、Stable Diffusion 3/XL、SD‑2、Z‑Image‑Turbo、CogVideoX‑2B、Wan‑2.1‑T2V‑1.3B、HunyuanVideo‑1.5 など。
シンプルな API 通常の Diffusers パイプライン呼び出しの周囲にコンテキストマネージャ AttentionCapture(pipe, relations=…, offload=…) を使用します。
自動可視化 キャプチャ後、compute().save(...) で PNG/GIF オーバーレイと JSON メタデータファイルを出力します。
動画処理対応 動画モデルでは、(T, H, W) パッチグリッドを保持します。フレームごとの PNG およびアニメーション GIF をレンダリングできます。
リソースに配慮した設計 デフォルトで 8 GiB のガードにより RAM/VRAM のオーバーフローを防止。max_capture_bytes を調整し、記録するタイムステップを選択できます。
CLI デモと検証 demo/run_attention_demo.py で任意の対応モデルの簡単な例を実行。テストスイート(pytest)および監査スクリプトでカバレッジと元の生成との整合性を検証。
インストール 1 つの pip コマンドでインストール:pip install attention-map-diffusers==1.0.0
ライセンスと引用 MIT ライセンス。学術利用のための DOI と BibTeX エントリを提供。

使い始め方(README からのクイックスタートスニペット)

画像例(FLUX.2‑Klein)

import torch
from diffusers import Flux2KleinPipeline
from attention_map_diffusers import AttentionCapture, text_tokenizers

prompt = "A red fox beside a glowing blue lantern in a snowy forest."
pipe = Flux2KleinPipeline.from_pretrained(
    "black-forest-labs/FLUX.2-klein-4B", torch_dtype=torch.bfloat16
).to("cuda")

with AttentionCapture(
    pipe,
    relations=["text->text", "text->image", "image->text", "image->image"],
    offload="cuda",
) as capture:
    images = pipe(prompt=[prompt], num_inference_steps=4).images

capture.compute().save(
    "outputs/attention", tokenizer=text_tokenizers(pipe),
    prompts=[prompt], images=images,
)

動画例(Wan‑2.1‑T2V‑1.3B)

import torch
from diffusers import WanPipeline
from attention_map_diffusers import AttentionCapture, VisualizationConfig, text_tokenizers

prompt = "A cinematic tracking shot of a red fox running across snow in a pine forest."
pipe = WanPipeline.from_pretrained(
    "Wan-AI/Wan2.1-T2V-1.3B-Diffusers", torch_dtype=torch.bfloat16
).to("cuda")
pipe.scheduler.set_timesteps(50, device="cuda")
last_timestep = [float(pipe.scheduler.timesteps[-1])]

with AttentionCapture(
    pipe,
    relations=["video->text", "video->video"],
    timesteps=last_timestep,
    relation_query_indices={"video->video": "center"},
    offload="cpu",
    max_capture_bytes=16*1024**3,
) as capture:
    videos = pipe(
        prompt=[prompt], num_inference_steps=50,
        height=480, width=832, num_frames=81,
    ).frames

# モジュールを CPU に戻して GPU メモリを解放
for comp in pipe.components.values():
    if isinstance(comp, torch.nn.Module):
        comp.to("cpu")
torch.cuda.empty_cache()

capture.compute(compute_device="cuda").save(
    "outputs/attention", tokenizer=text_tokenizers(pipe),
    prompts=[prompt], videos=videos,
    visualization_config=VisualizationConfig(max_items=16, video_fps=16),
)

ディスク上に得られるもの

このパッケージは以下のディレクトリ構造を出力します:

attention/
  metadata.json                # キャプチャ設定および出典情報
  raw/<relation>/*.pt          # --save-raw オプションを使用した場合にのみ保存される生のテンソル
  visuals/<relation>/aggregate/
    maps/*.png                 # フレームごとまたは画像ごとのヒートマップ
    maps/*.gif                 # 動画関係用のアニメーション GIF
    overlays/…                 # 元のメディアにアテンションをオーバーレイした画像

メタデータにはモデルのチェックポイント、プロンプト、キャプチャされたタイムステップ、および適用されたリソース制限が含まれます。

どのような人が使うか

  • 研究者:拡散モデルがプロンプトトークンや画像パッチにどのように注目しているかを調査する。
  • プロンプトエンジニア:特定のフレーズが出力のどの領域に影響を与えているかを確認したい人。
  • 教育者:生成AIにおけるクロスアテンションの内部動作を説明する。
  • 開発者:Diffusers パイプライン向けのデバッグツールや視覚的説明拡張機能を開発する。

  • 上記のすべての情報はリポジトリの README から直接取得したものであり、追加の機能は推測されていません。*

関連