Download memory_clip.py from AbstractPhil/geolip-clip-vit-large-patch14-ctx576: direct link, hf CLI and curl.
- Browser
- Download file 16.8 kB
-
https://huggingface.co/AbstractPhil/geolip-clip-vit-large-patch14-ctx576/resolve/main/memory_clip.py
- Command line
-
hf download hf://AbstractPhil/geolip-clip-vit-large-patch14-ctx576/memory_clip.py
-
curl -L -o memory_clip.py https://huggingface.co/AbstractPhil/geolip-clip-vit-large-patch14-ctx576/resolve/main/memory_clip.py
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) | |
| def n_extract_layers(self): | |
| return len(self.extract_layers) | |
| def depth_profile_dim(self): | |
| return self.n_extract_layers * self.clip_hidden | |
| 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() | |
| 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 | |
| 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 |