""" 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