AbstractPhil's picture
Rename memory_model_code.py to memory_clip.py
039183a verified
Raw History Blame Contribute Delete
16.8 kB
"""
MemoryCLIP — Memory-Extended CLIP-L/14 Text Encoder
Single file containing both config and model for HuggingFace AutoModel.
No cross-file imports.
Usage:
from transformers import AutoModel, AutoConfig
model = AutoModel.from_pretrained(
"AbstractPhil/geolip-clip-vit-large-patch14-ctx576",
trust_remote_code=True)
emb = model.encode("A long detailed caption...")
"""
import math
from typing import Optional, List
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
from transformers import PretrainedConfig, PreTrainedModel, CLIPTextModel, CLIPTokenizer
from transformers.modeling_outputs import BaseModelOutput
# ══════════════════════════════════════════════════════════════════
# CONFIG
# ══════════════════════════════════════════════════════════════════
class MemoryCLIPConfig(PretrainedConfig):
model_type = "memory_clip"
def __init__(
self,
clip_model="openai/clip-vit-large-patch14",
clip_hidden=768,
clip_layers=12,
clip_max_tokens=77,
freeze_clip=True,
n_memory_tokens=8,
bank_size=64,
anchor_dim=768,
n_bank_heads=8,
bank_cross_layers=2,
gate_type="gru",
extract_layers=(1, 3, 5, 7, 9, 11),
layer_fusion="learned",
max_content_tokens=18,
segment_overlap=4,
max_segments=32,
cv_target=0.20,
**kwargs,
):
self.clip_model = clip_model
self.clip_hidden = clip_hidden
self.clip_layers = clip_layers
self.clip_max_tokens = clip_max_tokens
self.freeze_clip = freeze_clip
self.n_memory_tokens = n_memory_tokens
self.bank_size = bank_size
self.anchor_dim = anchor_dim
self.n_bank_heads = n_bank_heads
self.bank_cross_layers = bank_cross_layers
self.gate_type = gate_type
self.extract_layers = tuple(extract_layers)
self.layer_fusion = layer_fusion
self.max_content_tokens = max_content_tokens
self.segment_overlap = segment_overlap
self.max_segments = max_segments
self.cv_target = cv_target
super().__init__(**kwargs)
@property
def n_extract_layers(self):
return len(self.extract_layers)
@property
def depth_profile_dim(self):
return self.n_extract_layers * self.clip_hidden
@property
def effective_context(self):
return self.max_segments * self.max_content_tokens
# ══════════════════════════════════════════════════════════════════
# COMPONENTS
# ══════════════════════════════════════════════════════════════════
class GeometricMemoryBank(nn.Module):
def __init__(self, config):
super().__init__()
self.max_size = config.bank_size
self.dim = config.anchor_dim
self.depth_compressor = nn.Sequential(
nn.Linear(config.depth_profile_dim, config.clip_hidden * 2),
nn.GELU(),
nn.LayerNorm(config.clip_hidden * 2),
nn.Linear(config.clip_hidden * 2, config.anchor_dim),
)
self.temporal_proj = nn.Linear(1, config.anchor_dim, bias=False)
self.cross_attn = nn.ModuleList([
nn.MultiheadAttention(config.clip_hidden, config.n_bank_heads,
batch_first=True, dropout=0.1)
for _ in range(config.bank_cross_layers)
])
self.cross_norms = nn.ModuleList([
nn.LayerNorm(config.clip_hidden)
for _ in range(config.bank_cross_layers)
])
self.cross_ffns = nn.ModuleList([
nn.Sequential(
nn.Linear(config.clip_hidden, config.clip_hidden * 2),
nn.GELU(),
nn.Linear(config.clip_hidden * 2, config.clip_hidden))
for _ in range(config.bank_cross_layers)
])
self.ffn_norms = nn.ModuleList([
nn.LayerNorm(config.clip_hidden)
for _ in range(config.bank_cross_layers)
])
def init_bank(self, batch_size, device):
return {"anchors": torch.zeros(batch_size, 0, self.dim, device=device),
"n_written": 0}
def write(self, bank, depth_cls, segment_idx=0):
B = depth_cls.shape[0]
anchor = self.depth_compressor(depth_cls.reshape(B, -1))
anchor = F.normalize(anchor, dim=-1)
t = torch.tensor([[segment_idx]], dtype=anchor.dtype, device=anchor.device)
anchor = anchor + 0.1 * self.temporal_proj(t / max(self.max_size, 1))
anchor = F.normalize(anchor, dim=-1)
anchors = torch.cat([bank["anchors"], anchor.unsqueeze(1)], dim=1)
if anchors.shape[1] > self.max_size:
anchors = anchors[:, -self.max_size:]
return {"anchors": anchors, "n_written": bank["n_written"] + 1,
"live_anchor": anchor}
def read(self, memory_tokens, bank):
anchors = bank["anchors"]
if anchors.shape[1] == 0:
return memory_tokens
x = memory_tokens
for attn, norm, ffn, ffn_norm in zip(
self.cross_attn, self.cross_norms, self.cross_ffns, self.ffn_norms):
residual = x
x, _ = attn(norm(x), anchors, anchors)
x = residual + x
residual = x
x = residual + ffn(ffn_norm(x))
return x
class DeltaMemoryGate(nn.Module):
def __init__(self, config):
super().__init__()
H = config.clip_hidden
self.reset_proj = nn.Linear(H * 2, H)
self.update_proj = nn.Linear(H * 2, H)
self.candidate_proj = nn.Linear(H * 2, H)
self.norm = nn.LayerNorm(H)
def forward(self, old, new):
cat = torch.cat([old, new], dim=-1)
r = torch.sigmoid(self.reset_proj(cat))
z = torch.sigmoid(self.update_proj(cat))
h = torch.tanh(self.candidate_proj(torch.cat([r * old, new], dim=-1)))
return self.norm(z * old + (1 - z) * h)
class LayerFusion(nn.Module):
def __init__(self, config):
super().__init__()
n = config.n_extract_layers
self.weights = nn.Parameter(torch.ones(n) / n)
self.proj = nn.Linear(config.clip_hidden, config.clip_hidden)
self.norm = nn.LayerNorm(config.clip_hidden)
def forward(self, layer_outputs):
w = F.softmax(self.weights, dim=0)
stacked = torch.stack(layer_outputs)
fused = (stacked * w.view(-1, 1, 1, 1)).sum(0)
return self.norm(self.proj(fused))
class TeacherProjector(nn.Module):
def __init__(self, student_dim, teacher_dim):
super().__init__()
self.proj = nn.Linear(student_dim, teacher_dim, bias=True)
def forward(self, x):
return self.proj(x)
# ══════════════════════════════════════════════════════════════════
# SEGMENTATION
# ══════════════════════════════════════════════════════════════════
def segment_text(text, clip_tokenizer, max_content=18, overlap=4, max_segments=32):
full_tokens = clip_tokenizer.encode(text, add_special_tokens=False)
segments = []
stride = max_content - overlap
pos = 0
while pos < len(full_tokens) and len(segments) < max_segments:
end = min(pos + max_content, len(full_tokens))
chunk = full_tokens[pos:end]
sos = clip_tokenizer.bos_token_id or 49406
eos = clip_tokenizer.eos_token_id or 49407
input_ids = [sos] + chunk + [eos]
n_pad = 77 - len(input_ids)
if n_pad > 0:
input_ids = input_ids + [0] * n_pad
else:
input_ids = input_ids[:77]
mask = [1] * min(len(chunk) + 2, 77) + [0] * max(n_pad, 0)
mask = mask[:77]
segments.append({
"input_ids": torch.tensor(input_ids, dtype=torch.long),
"attention_mask": torch.tensor(mask, dtype=torch.long),
})
if end >= len(full_tokens):
break
pos += stride
return segments
# ══════════════════════════════════════════════════════════════════
# MODEL
# ══════════════════════════════════════════════════════════════════
class MemoryCLIPModel(PreTrainedModel):
"""
Memory-Extended CLIP-L/14 Text Encoder.
Extends CLIP's 77-token context to 576 effective tokens via
geometric memory bank with depth-profile anchors.
"""
config_class = MemoryCLIPConfig
supports_gradient_checkpointing = False
def __init__(self, config):
super().__init__(config)
self.memory_embeddings = nn.Parameter(
torch.randn(1, config.n_memory_tokens, config.clip_hidden) * 0.02)
self.layer_fusion = LayerFusion(config)
self.bank = GeometricMemoryBank(config)
self.gate = DeltaMemoryGate(config)
self.output_proj = nn.Sequential(
nn.Linear(config.clip_hidden, config.clip_hidden),
nn.GELU(), nn.LayerNorm(config.clip_hidden))
self.memory_output_fusion = nn.Sequential(
nn.Linear(config.clip_hidden * 2, config.clip_hidden),
nn.GELU(),
nn.Linear(config.clip_hidden, config.clip_hidden))
self.clip_cross_attn = nn.ModuleList([
nn.MultiheadAttention(config.clip_hidden, config.n_bank_heads,
batch_first=True, dropout=0.1)
for _ in range(config.bank_cross_layers)
])
self.clip_cross_norms = nn.ModuleList([
nn.LayerNorm(config.clip_hidden)
for _ in range(config.bank_cross_layers)
])
self.clip_cross_ffns = nn.ModuleList([
nn.Sequential(
nn.Linear(config.clip_hidden, config.clip_hidden * 2),
nn.GELU(),
nn.Linear(config.clip_hidden * 2, config.clip_hidden))
for _ in range(config.bank_cross_layers)
])
self.clip_cross_ffn_norms = nn.ModuleList([
nn.LayerNorm(config.clip_hidden)
for _ in range(config.bank_cross_layers)
])
self.proj_modern = TeacherProjector(config.clip_hidden, 1024)
self._clip_text = None
self._clip_tokenizer = None
self.post_init()
@property
def clip_text(self):
if self._clip_text is None:
self._clip_text = CLIPTextModel.from_pretrained(self.config.clip_model)
self._clip_text.config.output_hidden_states = True
for p in self._clip_text.parameters():
p.requires_grad = False
device = self.memory_embeddings.device
self._clip_text = self._clip_text.to(device)
return self._clip_text
@property
def clip_tokenizer(self):
if self._clip_tokenizer is None:
self._clip_tokenizer = CLIPTokenizer.from_pretrained(self.config.clip_model)
return self._clip_tokenizer
def init_state(self, batch_size, device=None):
if device is None:
device = self.memory_embeddings.device
return {
"memory": self.memory_embeddings.expand(batch_size, -1, -1).clone(),
"bank": self.bank.init_bank(batch_size, device),
"segment_idx": 0,
}
def forward_segment(self, input_ids, attention_mask, state):
B = input_ids.shape[0]
memory_state = state["memory"]
bank = state["bank"]
seg_idx = state["segment_idx"]
memory_tokens = self.bank.read(memory_state, bank)
max_len = self.config.clip_max_tokens
with torch.no_grad():
clip_out = self.clip_text(
input_ids=input_ids[:, :max_len],
attention_mask=attention_mask[:, :max_len],
output_hidden_states=True, return_dict=True)
all_hiddens = clip_out.hidden_states
selected = [all_hiddens[i + 1] for i in self.config.extract_layers]
fused = self.layer_fusion(selected)
mem_enriched = memory_tokens
for attn, norm, ffn, ffn_norm in zip(
self.clip_cross_attn, self.clip_cross_norms,
self.clip_cross_ffns, self.clip_cross_ffn_norms):
residual = mem_enriched
mem_enriched, _ = attn(norm(mem_enriched), fused, fused)
mem_enriched = residual + mem_enriched
residual = mem_enriched
mem_enriched = residual + ffn(ffn_norm(mem_enriched))
depth_cls = torch.stack([h[:, 1, :] for h in selected], dim=1)
new_memory = self.gate(memory_state, mem_enriched)
new_bank = self.bank.write(bank, depth_cls, seg_idx)
clip_pooled = clip_out.pooler_output
if clip_pooled is None:
clip_pooled = clip_out.last_hidden_state[:, -1, :]
cls_output = self.output_proj(clip_pooled)
memory_delta = self.memory_output_fusion(
torch.cat([cls_output, new_memory.mean(dim=1)], dim=-1))
fused_output = cls_output + memory_delta
new_state = {
"memory": new_memory,
"bank": {"anchors": new_bank["anchors"],
"n_written": new_bank["n_written"]},
"segment_idx": seg_idx + 1,
}
return fused_output, new_state
def forward(
self,
input_ids: Optional[torch.LongTensor] = None,
attention_mask: Optional[torch.Tensor] = None,
texts: Optional[List[str]] = None,
return_dict: bool = True,
**kwargs,
) -> BaseModelOutput:
"""
Accepts either:
- input_ids + attention_mask (single 77-token segment)
- texts (list of strings, auto-segmented for long context)
Returns BaseModelOutput with last_hidden_state = (B, 1, 768).
"""
device = self.memory_embeddings.device
if texts is not None:
embeddings = [self._encode_single(t, device) for t in texts]
last_hidden = torch.stack(embeddings)
elif input_ids is not None:
state = self.init_state(input_ids.shape[0], device)
last_hidden, _ = self.forward_segment(
input_ids.to(device), attention_mask.to(device), state)
else:
raise ValueError("Provide either input_ids or texts")
if return_dict:
return BaseModelOutput(
last_hidden_state=last_hidden.unsqueeze(1),
hidden_states=None, attentions=None)
return (last_hidden.unsqueeze(1),)
def _encode_single(self, text, device):
segments = segment_text(
text, self.clip_tokenizer,
self.config.max_content_tokens,
self.config.segment_overlap,
self.config.max_segments)
state = self.init_state(1, device)
output = None
for seg in segments:
ids = seg["input_ids"].unsqueeze(0).to(device)
mask = seg["attention_mask"].unsqueeze(0).to(device)
output, state = self.forward_segment(ids, mask, state)
return output.squeeze(0)
def encode(self, texts, batch_size=32, show_progress=False):
"""
Encode text(s) to 768-dim CLIP-compatible embeddings.
Returns:
torch.Tensor: (N, 768) or (768,) for single string
"""
device = self.memory_embeddings.device
single = isinstance(texts, str)
if single:
texts = [texts]
all_embs = []
iterator = range(0, len(texts), batch_size)
if show_progress:
from tqdm import tqdm
iterator = tqdm(iterator, desc="Encoding")
with torch.no_grad():
for i in iterator:
batch = texts[i:i + batch_size]
embs = [self._encode_single(t, device) for t in batch]
all_embs.append(torch.stack(embs))
result = torch.cat(all_embs, dim=0)
return result.squeeze(0) if single else result