Mistral: use fork-capable attention bridge (enable hook_attn_in / set_use_attn_in)#1497
Conversation
Mistral has separate q/k/v/o projections with RoPE and GQA, structurally
identical to Qwen2, but its adapter used the plain AttentionBridge, which
delegates q/k/v to HF and exposes no per-receiver fork point. As a result
set_use_attn_in() raised and blocks.{i}.hook_attn_in did not exist, so
per-head input interventions (activation/path patching on attention inputs)
were impossible on Mistral.
Switch the Mistral adapter to PositionEmbeddingsAttentionBridge (the same
fork-capable bridge Qwen2/Qwen3 use) and pass config to its BlockBridge so
hook_mlp_in is wired too. Add a test booting the tiny random Mistral that
checks the attention bridge is fork-capable, hook_attn_in fires at
[batch, pos, n_heads, d_model], and zeroing a single head's forked input
actually changes the logits.
|
@almogtavor This looks good to me, from my review. Please just update or remove the stale test that is causing CI to fail. |
|
@almogtavor I did find one additional item while running some Run |
|
@jlarson4 Awesome, I'll fix it thanks ! |
|
Hey @almogtavor! Just wanted to touch base and see if you'll have an opportunity to wrap this up within the next couple days. I am getting ready to do a release, and would love to include this if it will be ready. Thank you! |
|
@jlarson4 Of course, I'll wrap it up on the weekend :) |
PositionEmbeddingsAttentionBridge used the first param's dtype, which is integer storage on quantized models (uint8 bnb, int32 GPTQ), so fp hidden states got cast to it. Ported the guard from AttentionBridge. Also updates the two Mistral tests that still assumed a plain AttentionBridge. The quantized one reads the dtype at q_proj now, since the reimplemented bridge never calls the HF attention forward.
|
@jlarson4 Both fixed :) The stale assert is inverted now. For quantized, ported the non-fp param skip into |
|
Looks great on my end! Thanks for putting this together! Merging now @almogtavor |
Mistral loads through the plain
AttentionBridge, which delegates q/k/v to HF and exposes no per-receiver fork point. So on a bridged Mistral,set_use_attn_in(True)raises andblocks.{i}.hook_attn_in(and the q/k/v input forks) never exist, making per-head attention-input interventions (activation/path patching that edits a receiver input) impossible.Mistral is structurally identical to Qwen2 here (separate
q/k/v/o_proj, RoPE, GQA, RMSNorm, gated MLP), which usesPositionEmbeddingsAttentionBridge.Fix
PositionEmbeddingsAttentionBridge(same bridge as Qwen2/Qwen3).configto Mistral'sBlockBridgesohook_mlp_inis wired too.Test
tests/unit/model_bridge/test_mistral_attn_in_fork.pybootshf-internal-testing/tiny-random-MistralForCausalLMand checks the bridge isPositionEmbeddingsAttentionBridge,hook_attn_infires at[batch, pos, n_heads, d_model], and zeroing a single head's forked input changes the logits.