ToiTenBao's picture
Upload hallucination folder
a2ffd07 verified
Raw
History Blame Contribute Delete
11.4 kB
"""
Cross-attention adapter modules for DualEdit.
Adapted from DualEdit/editor/vllm_editors/vead/adpt_model.py.
Two adapter types:
- VisionEditAdapter: modifies visual token representations
- TextEditAdapter: modifies text token representations
Each uses cross-attention between current hidden states and cached edit signals
to inject edited knowledge at specific transformer layers.
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
class VisionEditAdapter(nn.Module):
"""Cross-attention adapter for editing visual token representations.
Inserted at a specific transformer layer via forward hook. Uses the edit
signal (cached hidden states from the edit sample) as Key/Value, and the
current image token hidden states as Query.
Args:
hidden_size: Model hidden dimension (4096 for LLaMA-7B).
mid_dim: Adapter bottleneck dimension.
cross_att_head_n: Number of cross-attention heads.
img_tok_n: Number of image tokens (576 for LLaVA-1.5).
"""
def __init__(self, hidden_size, mid_dim=1024, cross_att_head_n=8, img_tok_n=576):
super().__init__()
if mid_dim % cross_att_head_n != 0:
raise ValueError(f"mid_dim ({mid_dim}) must be divisible by cross_att_head_n ({cross_att_head_n})")
self.mid_dim = mid_dim
self.cross_att_head_n = cross_att_head_n
self.img_tok_n = img_tok_n
self.mlp_begin = nn.Linear(hidden_size, mid_dim)
self.cross_att_q_mlp = nn.Linear(mid_dim, mid_dim)
self.cross_att_k_mlp = nn.Linear(hidden_size, mid_dim)
self.cross_att_v_mlp = nn.Linear(hidden_size, mid_dim)
self.mlp_end = nn.Linear(mid_dim, hidden_size)
self.ln_img_reps = nn.LayerNorm(hidden_size)
self.ln_edit_reps = nn.LayerNorm(hidden_size)
# State
self.is_open = False
self.open_gating = False
self.edit_reps = None
self.edit_reps_att_mask = None
self.inpt_has_img = True
self.inpt_vt_begin = None
self.inpt_vt_end = None
# Gate prototype for inference-time gating
self.gate_prototype = None
self.gate_threshold = 0.6
def forward(self, layer_outpt):
"""Apply adapter to layer output.
Args:
layer_outpt: [batch, seq_len, hidden_size] tensor from transformer layer.
Returns:
Modified layer output with edited image token representations.
"""
if (not self.is_open
or layer_outpt.shape[1] == 1 # generation mode (kv cache)
or not self.inpt_has_img
or self.edit_reps is None):
return layer_outpt
orig_dtype = layer_outpt.dtype
# Upcast to float32 for numerical stability (matches original DualEdit)
layer_outpt = layer_outpt.float()
layer_input = layer_outpt.clone()
if self.inpt_vt_begin is None or self.inpt_vt_end is None:
return layer_outpt.to(orig_dtype)
# Extract image tokens
img_reps = layer_outpt[:, self.inpt_vt_begin:self.inpt_vt_end].clone()
b1, l1, _ = img_reps.shape
b2, l2, _ = self.edit_reps.shape
if l1 != self.img_tok_n:
return layer_outpt.to(orig_dtype)
if b1 != b2:
if b2 == 1:
edit_reps = self.edit_reps.float().expand(b1, -1, -1)
edit_mask = self.edit_reps_att_mask.float().expand(b1, -1)
else:
return layer_outpt.to(orig_dtype)
else:
edit_reps = self.edit_reps.float()
edit_mask = self.edit_reps_att_mask.float()
# Cross-attention: image tokens attend to edit signal
norm_img_reps = self.ln_img_reps(img_reps)
norm_edit_reps = self.ln_edit_reps(edit_reps)
x = self.mlp_begin(norm_img_reps)
q = self.cross_att_q_mlp(x).reshape(b1, l1, self.cross_att_head_n, self.mid_dim // self.cross_att_head_n)
k = self.cross_att_k_mlp(norm_edit_reps).reshape(b1, l2, self.cross_att_head_n, self.mid_dim // self.cross_att_head_n)
v = self.cross_att_v_mlp(norm_edit_reps).reshape(b1, l2, self.cross_att_head_n, self.mid_dim // self.cross_att_head_n)
s = torch.einsum('blhm,buhm->bhlu', q, k)
s = s / (self.mid_dim // self.cross_att_head_n) ** 0.5
s = s + (edit_mask.reshape(b1, 1, 1, l2) - 1) * 9999999999
s = torch.softmax(s, dim=3)
x = torch.einsum('bhlu,buhm->blhm', s, v).reshape(b1, l1, self.mid_dim)
x = self.mlp_end(x)
# Residual connection
layer_outpt[:, self.inpt_vt_begin:self.inpt_vt_end] = img_reps + x
# Gating: if enabled, decide per-sample whether to apply edit
if self.open_gating and self.gate_prototype is not None:
sim = F.cosine_similarity(
layer_outpt[:, -1, :],
self.gate_prototype.float().unsqueeze(0),
dim=-1,
)
should_edit = (sim > self.gate_threshold).unsqueeze(-1).unsqueeze(-1)
print(f" [DualEdit VisionAdapter] gate sim={sim.tolist()}, threshold={self.gate_threshold}, fires={should_edit.squeeze().tolist()}")
layer_outpt = torch.where(should_edit.expand_as(layer_outpt), layer_outpt, layer_input)
return layer_outpt.to(orig_dtype)
def open_adapter(self, is_open: bool):
self.is_open = is_open
def set_edit_signal(self, edit_reps, edit_reps_att_mask):
self.edit_reps = edit_reps
self.edit_reps_att_mask = edit_reps_att_mask
def set_input_info(self, has_img=True, vt_begin=None, vt_end=None):
self.inpt_has_img = has_img
self.inpt_vt_begin = vt_begin
self.inpt_vt_end = vt_end
def set_gate(self, prototype, threshold=0.6):
self.gate_prototype = prototype
self.gate_threshold = threshold
class TextEditAdapter(nn.Module):
"""Cross-attention adapter for editing text token representations.
Similar to VisionEditAdapter but operates on text tokens (non-image tokens).
Processes each sample separately due to variable text token counts.
Args:
hidden_size: Model hidden dimension.
mid_dim: Adapter bottleneck dimension.
cross_att_head_n: Number of cross-attention heads.
"""
def __init__(self, hidden_size, mid_dim=1024, cross_att_head_n=8):
super().__init__()
if mid_dim % cross_att_head_n != 0:
raise ValueError(f"mid_dim ({mid_dim}) must be divisible by cross_att_head_n ({cross_att_head_n})")
self.mid_dim = mid_dim
self.cross_att_head_n = cross_att_head_n
self.mlp_begin = nn.Linear(hidden_size, mid_dim)
self.cross_att_q_mlp = nn.Linear(mid_dim, mid_dim)
self.cross_att_k_mlp = nn.Linear(hidden_size, mid_dim)
self.cross_att_v_mlp = nn.Linear(hidden_size, mid_dim)
self.mlp_end = nn.Linear(mid_dim, hidden_size)
self.ln_text_reps = nn.LayerNorm(hidden_size)
self.ln_edit_reps = nn.LayerNorm(hidden_size)
# State
self.is_open = False
self.open_gating = False
self.edit_reps = None
self.edit_reps_att_mask = None
self.prompt_end = None
self.inpt_vt_end = None
# Gate prototype
self.gate_prototype = None
self.gate_threshold = 0.6
def forward(self, layer_outpt):
"""Apply adapter to layer output for text tokens."""
# Handle tuple output from transformer layers
is_tuple = isinstance(layer_outpt, tuple)
if is_tuple:
layer_outpt = list(layer_outpt)
hidden = layer_outpt[0]
else:
hidden = layer_outpt
if (not self.is_open
or hidden.shape[1] == 1 # generation mode
or self.edit_reps is None
or self.prompt_end is None):
return tuple(layer_outpt) if is_tuple else layer_outpt
orig_dtype = hidden.dtype
# Upcast to float32 for numerical stability
hidden = hidden.float()
layer_input = hidden.clone()
batch_size = hidden.shape[0]
for i in range(batch_size):
if self.inpt_vt_end is not None:
if self.prompt_end.dim() == 0:
end = self.prompt_end.item()
else:
end = self.prompt_end[i].item() if i < len(self.prompt_end) else hidden.shape[1]
indices = list(range(int(self.inpt_vt_end), int(end)))
else:
if self.prompt_end.dim() == 0:
end = self.prompt_end.item()
else:
end = self.prompt_end[i].item() if i < len(self.prompt_end) else hidden.shape[1]
indices = list(range(1, int(end)))
if not indices:
continue
text_reps = hidden[i, indices].unsqueeze(0) # [1, n_text, d]
b1, l1, _ = text_reps.shape
if self.edit_reps.shape[0] > 1 and i < self.edit_reps.shape[0]:
sample_edit = self.edit_reps[i:i + 1].float()
sample_mask = self.edit_reps_att_mask[i:i + 1].float()
else:
sample_edit = self.edit_reps[:1].float()
sample_mask = self.edit_reps_att_mask[:1].float()
l2 = sample_edit.shape[1]
norm_text = self.ln_text_reps(text_reps)
norm_edit = self.ln_edit_reps(sample_edit)
x = self.mlp_begin(norm_text)
q = self.cross_att_q_mlp(x).reshape(1, l1, self.cross_att_head_n, self.mid_dim // self.cross_att_head_n)
k = self.cross_att_k_mlp(norm_edit).reshape(1, l2, self.cross_att_head_n, self.mid_dim // self.cross_att_head_n)
v = self.cross_att_v_mlp(norm_edit).reshape(1, l2, self.cross_att_head_n, self.mid_dim // self.cross_att_head_n)
s = torch.einsum('blhm,buhm->bhlu', q, k)
s = s / (self.mid_dim // self.cross_att_head_n) ** 0.5
s = s + (sample_mask.reshape(1, 1, 1, l2) - 1) * 9999999999
s = torch.softmax(s, dim=3)
x = torch.einsum('bhlu,buhm->blhm', s, v).reshape(1, l1, self.mid_dim)
x = self.mlp_end(x)
hidden[i, indices] = text_reps.squeeze(0) + x.squeeze(0)
# Gating
if self.open_gating and self.gate_prototype is not None:
sim = F.cosine_similarity(
hidden[:, -1, :],
self.gate_prototype.float().unsqueeze(0),
dim=-1,
)
should_edit = (sim > self.gate_threshold).unsqueeze(-1).unsqueeze(-1)
hidden = torch.where(should_edit.expand_as(hidden), hidden, layer_input)
hidden = hidden.to(orig_dtype)
if is_tuple:
layer_outpt[0] = hidden
return tuple(layer_outpt)
return hidden
def open_adapter(self, is_open: bool):
self.is_open = is_open
def set_edit_signal(self, edit_reps, edit_reps_att_mask, prompt_end=None):
self.edit_reps = edit_reps
self.edit_reps_att_mask = edit_reps_att_mask
self.prompt_end = prompt_end
def set_input_info(self, has_img=True, vt_begin=None, vt_end=None):
self.inpt_vt_end = vt_end
def set_gate(self, prototype, threshold=0.6):
self.gate_prototype = prototype
self.gate_threshold = threshold