| """ |
| 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) |
|
|
| |
| 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 |
| |
| 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 |
| or not self.inpt_has_img |
| or self.edit_reps is None): |
| return layer_outpt |
|
|
| orig_dtype = layer_outpt.dtype |
| |
| 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) |
|
|
| |
| 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() |
|
|
| |
| 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) |
|
|
| |
| layer_outpt[:, self.inpt_vt_begin:self.inpt_vt_end] = img_reps + x |
|
|
| |
| 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) |
|
|
| |
| 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 |
| |
| self.gate_prototype = None |
| self.gate_threshold = 0.6 |
|
|
| def forward(self, layer_outpt): |
| """Apply adapter to layer output for text tokens.""" |
| |
| 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 |
| 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 |
| |
| 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) |
| 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) |
|
|
| |
| 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 |
|
|