Feature Extraction
Transformers
Safetensors
prism
video
representation-learning
view-invariant
cross-view
egocentric
egoexo4d
emnlp2026
custom_code
Instructions to use litcoderr/prism with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use litcoderr/prism with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("feature-extraction", model="litcoderr/prism", trust_remote_code=True)# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("litcoderr/prism", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
| """PRISM model. | |
| Frozen SigLIP2 vision + frozen Qwen3-Embedding text + trainable Decompositional | |
| Encoder θ + Compositional Latent Predictor φ + EMA target encoder θ̄. | |
| One training ``forward`` returns: | |
| - ``loss_decomp`` : symmetric InfoNCE between the compositional latent | |
| ``s = φ(z_vv^B, z_vi^A)`` and the recomposed text embedding | |
| ``e = Qwen3Embedding(compose(T_vi^A, T_vv^B))``, over the batch | |
| (with DDP all-gather of ``s`` / ``e`` / the valid-pair mask). [paper §3.2] | |
| - ``loss_temp_vi`` / ``loss_temp_vv`` : ``1 - cos(ẑ_t, z̄_{t+1})`` for each | |
| stream, where the target ``z̄`` comes from the EMA encoder θ̄. [paper §3.3] | |
| - ``loss = λ_decomp · loss_decomp + λ_temp · ½(loss_temp_vi + loss_temp_vv)``. | |
| Cross-pairing (one clip's view-variant stream with another clip's view-invariant | |
| stream) is done inside ``forward`` via a cyclic shift of the batch: the | |
| view-variant stream comes from clip ``i``, the view-invariant stream from clip | |
| ``(i+1) mod B``. The trainer builds the matching recomposed caption with the | |
| same convention. | |
| At inference, ``encode`` returns an L2-normalized clip embedding (mean-pooled | |
| ``z_vi`` over valid frames) from the EMA encoder. | |
| """ | |
| from __future__ import annotations | |
| import logging | |
| from dataclasses import dataclass | |
| from pathlib import Path | |
| import torch | |
| import torch.distributed as dist | |
| import torch.nn.functional as F | |
| from torch import nn | |
| from transformers import AutoModel, PreTrainedModel | |
| logger = logging.getLogger(__name__) | |
| from .configuration_prism import PRISMConfig | |
| from .ema import make_ema_copy, sync_ema_from_online, update_ema | |
| from .encoder import DecompositionalEncoder | |
| # Unused here, but kept as a direct import: when this file is served as Hub remote | |
| # code, transformers ships only the relative imports named in *this* module, so | |
| # ``layers`` (used by encoder/predictor) has to be visible from here. | |
| from .layers import QFormerBlock, TemporalBlock # noqa: F401 | |
| from .predictor import CompositionalPredictor | |
| class PRISMOutput: | |
| loss: torch.Tensor | |
| loss_decomp: torch.Tensor | |
| loss_temp_vi: torch.Tensor | |
| loss_temp_vv: torch.Tensor | |
| n_valid_pairs: int | |
| # --------------------------------------------------------------------------- | |
| # Loss / distributed helpers | |
| # --------------------------------------------------------------------------- | |
| def _symmetric_infonce( | |
| a: torch.Tensor, b: torch.Tensor, logit_scale: torch.Tensor | |
| ) -> torch.Tensor: | |
| """CLIP-style symmetric InfoNCE on L2-normalized features.""" | |
| a = F.normalize(a, dim=-1) | |
| b = F.normalize(b, dim=-1) | |
| scale = logit_scale.exp().clamp(max=100.0) | |
| logits = scale * a @ b.t() | |
| labels = torch.arange(a.shape[0], device=a.device) | |
| return 0.5 * (F.cross_entropy(logits, labels) + F.cross_entropy(logits.t(), labels)) | |
| def _all_gather_with_grad(x: torch.Tensor) -> torch.Tensor: | |
| """CLIP-style all-gather: concat across ranks; own-rank slot keeps gradient. | |
| Other ranks' tensors are detached for this rank's backward; DDP's gradient | |
| all-reduce then distributes the gradient across ranks, making it equivalent | |
| to a single forward over the full ``B * world_size`` batch. Single-process | |
| → returns ``x`` unchanged. | |
| """ | |
| if not dist.is_available() or not dist.is_initialized(): | |
| return x | |
| world_size = dist.get_world_size() | |
| if world_size == 1: | |
| return x | |
| rank = dist.get_rank() | |
| gathered = [torch.empty_like(x) for _ in range(world_size)] | |
| dist.all_gather(gathered, x.contiguous()) | |
| gathered[rank] = x # own-rank slot keeps grad | |
| return torch.cat(gathered, dim=0) | |
| def _all_gather_bool(x: torch.Tensor) -> torch.Tensor: | |
| """Plain all-gather for boolean masks (no gradient).""" | |
| if not dist.is_available() or not dist.is_initialized(): | |
| return x | |
| world_size = dist.get_world_size() | |
| if world_size == 1: | |
| return x | |
| gathered = [torch.empty_like(x) for _ in range(world_size)] | |
| dist.all_gather(gathered, x.contiguous()) | |
| return torch.cat(gathered, dim=0) | |
| def _sample_shift_plan( | |
| valid_a: torch.Tensor, valid_b: torch.Tensor, T: int | |
| ) -> dict: | |
| """Sliding-shift augmentation plan. | |
| For each sample, place the shorter clip's valid frames at a random offset | |
| within the longer clip's valid range. Returns gather indices and post-shift | |
| valid masks for both sides; apply identically to online and EMA tensors so | |
| predictor input and temporal target stay time-aligned. | |
| """ | |
| B = valid_a.shape[0] | |
| device = valid_a.device | |
| t_a = valid_a.int().sum(dim=1) | |
| t_b = valid_b.int().sum(dim=1) | |
| max_off_a = torch.clamp(t_b - t_a, min=0) | |
| max_off_b = torch.clamp(t_a - t_b, min=0) | |
| rand = torch.rand(B, 2, device=device) | |
| offset_a = (rand[:, 0] * (max_off_a.float() + 1.0)).long().clamp(max=max_off_a) | |
| offset_b = (rand[:, 1] * (max_off_b.float() + 1.0)).long().clamp(max=max_off_b) | |
| arange_T = torch.arange(T, device=device).unsqueeze(0).expand(B, -1) | |
| src_idx_a = arange_T - offset_a.unsqueeze(1) | |
| src_idx_b = arange_T - offset_b.unsqueeze(1) | |
| new_valid_a = (src_idx_a >= 0) & (src_idx_a < t_a.unsqueeze(1)) | |
| new_valid_b = (src_idx_b >= 0) & (src_idx_b < t_b.unsqueeze(1)) | |
| return { | |
| "src_idx_a": src_idx_a.clamp(0, T - 1), | |
| "src_idx_b": src_idx_b.clamp(0, T - 1), | |
| "new_valid_a": new_valid_a, | |
| "new_valid_b": new_valid_b, | |
| } | |
| def _apply_shift( | |
| z: torch.Tensor, src_idx: torch.Tensor, valid: torch.Tensor | |
| ) -> torch.Tensor: | |
| """Gather ``z[B, T, D]`` along T via ``src_idx[B, T]``; zero invalid positions.""" | |
| gather_idx = src_idx.unsqueeze(-1).expand(-1, -1, z.shape[-1]) | |
| return torch.gather(z, dim=1, index=gather_idx) * valid.unsqueeze(-1).to(z.dtype) | |
| # --------------------------------------------------------------------------- | |
| # Checkpoint resolution (local directory or Hugging Face Hub repo) | |
| # --------------------------------------------------------------------------- | |
| _WEIGHTS_NAME = "model.safetensors" | |
| _HUB_KWARGS = ("revision", "cache_dir", "token", "force_download", "local_files_only", "proxies") | |
| def _resolve_weights(path_or_repo: str, **hub_kwargs) -> str: | |
| """Path to the checkpoint's weights: a local directory, else a Hub repo id.""" | |
| local = Path(path_or_repo) / _WEIGHTS_NAME | |
| if local.is_file(): | |
| return str(local) | |
| from huggingface_hub import hf_hub_download | |
| return hf_hub_download(repo_id=str(path_or_repo), filename=_WEIGHTS_NAME, **hub_kwargs) | |
| # --------------------------------------------------------------------------- | |
| # Model | |
| # --------------------------------------------------------------------------- | |
| class PRISMModel(PreTrainedModel): | |
| config_class = PRISMConfig | |
| base_model_prefix = "prism" | |
| # Frozen backbones are reloaded from the Hub in __init__ and excluded from | |
| # the saved checkpoint (see ``state_dict``); silence the load-time warning. | |
| _keys_to_ignore_on_load_missing = [r"^vision_model\.", r"^text_model\."] | |
| supports_gradient_checkpointing = False | |
| def __init__(self, config: PRISMConfig): | |
| super().__init__(config) | |
| # ---- Frozen vision tower ---- | |
| # CLIP keeps a CLS token in last_hidden_state; SigLIP / SigLIP2 do not. | |
| # SigLIP2 weights use the SigLIP v1 architecture, so SiglipVisionModel | |
| # handles both checkpoint families. | |
| if "siglip" in config.vision_backbone_name.lower(): | |
| from transformers import SiglipVisionModel | |
| self.vision_model = SiglipVisionModel.from_pretrained(config.vision_backbone_name) | |
| self._vision_has_cls = False | |
| else: | |
| from transformers import CLIPVisionModel | |
| self.vision_model = CLIPVisionModel.from_pretrained(config.vision_backbone_name) | |
| self._vision_has_cls = True | |
| d_v = int(self.vision_model.config.hidden_size) | |
| # ---- Frozen Qwen3-Embedding text tower ---- | |
| self.text_model = AutoModel.from_pretrained(config.text_backbone_name) | |
| d_t = int(self.text_model.config.hidden_size) | |
| for p in self.vision_model.parameters(): | |
| p.requires_grad = False | |
| for p in self.text_model.parameters(): | |
| p.requires_grad = False | |
| self.vision_model.eval() | |
| self.text_model.eval() | |
| self.d_v = d_v | |
| self.d_t = d_t | |
| # ---- Trainable: Decompositional Encoder θ + Compositional Predictor φ ---- | |
| self.encoder = DecompositionalEncoder( | |
| d_z=config.d_z, d_kv=d_v, | |
| qformer_depth=config.qformer_depth, temporal_depth=config.temporal_depth, | |
| num_heads=config.num_heads, mlp_ratio=config.mlp_ratio, | |
| max_frames=config.max_frames, | |
| ) | |
| self.predictor = CompositionalPredictor( | |
| d_z=config.d_z, d_t=d_t, max_frames=config.max_frames, | |
| depth=config.predictor_depth, num_heads=config.num_heads, | |
| mlp_ratio=config.mlp_ratio, | |
| ) | |
| # ---- EMA target encoder θ̄ ---- | |
| if config.use_ema: | |
| self.target_encoder = make_ema_copy(self.encoder) | |
| self.logit_scale = nn.Parameter(torch.tensor(config.logit_scale_init)) | |
| # -- keep trainable defaults; do not re-init the from-Hub backbones -- | |
| def _init_weights(self, module): # noqa: D401 | |
| pass | |
| def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs): | |
| """Load a weights-only PRISM checkpoint from a local directory or the Hub. | |
| The frozen vision / text backbones are not stored in the checkpoint; | |
| they are rebuilt from the Hub in ``__init__``. Only the trained weights | |
| (encoder θ, predictor φ, target encoder θ̄, logit scale) are loaded. The | |
| ``dtype`` / ``torch_dtype`` kwarg is honored; other HF loading kwargs | |
| (device_map, sharding, ...) are not needed for this single-file ckpt. | |
| """ | |
| from safetensors.torch import load_file | |
| hub_kwargs = {k: kwargs.pop(k) for k in _HUB_KWARGS if kwargs.get(k) is not None} | |
| config = kwargs.pop("config", None) | |
| if not isinstance(config, PRISMConfig): | |
| config = PRISMConfig.from_pretrained(pretrained_model_name_or_path, **hub_kwargs) | |
| dtype = kwargs.pop("torch_dtype", None) or kwargs.pop("dtype", None) | |
| model = cls(config) # backbones materialized from the Hub | |
| state = load_file(_resolve_weights(pretrained_model_name_or_path, **hub_kwargs)) | |
| missing, unexpected = model.load_state_dict(state, strict=False) | |
| bad_missing = [m for m in missing if not m.startswith(("vision_model.", "text_model."))] | |
| if bad_missing: | |
| logger.warning(f"missing non-backbone keys: {bad_missing[:8]}") | |
| if unexpected: | |
| logger.warning(f"unexpected keys: {unexpected[:8]}") | |
| # Warm-start the target encoder if a checkpoint predates EMA weights. | |
| if config.use_ema and any(k.startswith("target_encoder.") for k in bad_missing): | |
| model.sync_ema_from_online() | |
| if dtype is not None: | |
| model = model.to(dtype) | |
| return model | |
| def train(self, mode: bool = True): | |
| """Keep frozen backbones (and the EMA target encoder) in eval mode.""" | |
| super().train(mode) | |
| self.vision_model.eval() | |
| self.text_model.eval() | |
| if getattr(self.config, "use_ema", False): | |
| self.target_encoder.eval() | |
| return self | |
| def state_dict(self, *args, **kwargs): | |
| """Exclude the frozen, from-Hub backbones from saved checkpoints.""" | |
| sd = super().state_dict(*args, **kwargs) | |
| return type(sd)( | |
| (k, v) for k, v in sd.items() | |
| if not k.startswith(("vision_model.", "text_model.")) | |
| ) | |
| # ------------------------------------------------------------------ | |
| # Frozen backbone helpers | |
| # ------------------------------------------------------------------ | |
| def _encode_video(self, pixel_values: torch.Tensor) -> torch.Tensor: | |
| """``(B, T, 3, H, W)`` → patch tokens ``(B, T, P, d_v)`` (CLS dropped for CLIP).""" | |
| B, T = pixel_values.shape[:2] | |
| x = pixel_values.reshape(B * T, *pixel_values.shape[2:]) | |
| seq = self.vision_model(pixel_values=x).last_hidden_state | |
| if self._vision_has_cls: | |
| seq = seq[:, 1:, :] | |
| return seq.reshape(B, T, seq.shape[1], self.d_v) | |
| def _encode_text( | |
| self, input_ids: torch.Tensor, attention_mask: torch.Tensor | |
| ) -> torch.Tensor: | |
| """Qwen3-Embedding: last-token pool over a right-padded batch → L2-normed ``(N, d_t)``.""" | |
| last_hidden = self.text_model( | |
| input_ids=input_ids, attention_mask=attention_mask | |
| ).last_hidden_state | |
| last_idx = (attention_mask.sum(dim=1) - 1).clamp(min=0) | |
| rows = torch.arange(last_hidden.shape[0], device=last_hidden.device) | |
| return F.normalize(last_hidden[rows, last_idx], dim=-1) | |
| # ------------------------------------------------------------------ | |
| # EMA hooks (called by the trainer after each optimizer step) | |
| # ------------------------------------------------------------------ | |
| def update_ema(self) -> None: | |
| if getattr(self.config, "use_ema", False): | |
| update_ema(self.target_encoder, self.encoder, self.config.ema_decay) | |
| def sync_ema_from_online(self) -> None: | |
| if getattr(self.config, "use_ema", False): | |
| sync_ema_from_online(self.target_encoder, self.encoder) | |
| # ------------------------------------------------------------------ | |
| # Inference | |
| # ------------------------------------------------------------------ | |
| def encode_streams( | |
| self, pixel_values: torch.Tensor, valid_mask: torch.Tensor | None = None | |
| ) -> dict: | |
| """Run θ̄ (or θ if ``use_ema=False``) → ``{z_vi_seq, z_vv_seq}`` each ``(B, T, d_z)``.""" | |
| patches = self._encode_video(pixel_values) | |
| kpm = None if valid_mask is None else (~valid_mask) | |
| enc = self.target_encoder if getattr(self.config, "use_ema", False) else self.encoder | |
| z_vi, z_vv = enc(patches, key_padding_mask=kpm) | |
| return {"z_vi_seq": z_vi, "z_vv_seq": z_vv} | |
| def encode( | |
| self, pixel_values: torch.Tensor, valid_mask: torch.Tensor | None = None | |
| ) -> torch.Tensor: | |
| """Clip embedding: L2-normalized mean-pool of ``z_vi`` over valid frames → ``(B, d_z)``.""" | |
| z_vi = self.encode_streams(pixel_values, valid_mask)["z_vi_seq"].float() | |
| if valid_mask is None: | |
| valid = torch.ones(z_vi.shape[:2], device=z_vi.device, dtype=z_vi.dtype) | |
| else: | |
| valid = valid_mask.to(z_vi.dtype) | |
| denom = valid.sum(dim=1, keepdim=True).clamp(min=1.0) | |
| emb = (z_vi * valid.unsqueeze(-1)).sum(dim=1) / denom | |
| return F.normalize(emb, dim=-1) | |
| # ------------------------------------------------------------------ | |
| # Training forward | |
| # ------------------------------------------------------------------ | |
| def forward( | |
| self, | |
| pixel_values: torch.Tensor, | |
| valid_mask: torch.Tensor, | |
| composed_input_ids: torch.Tensor, | |
| composed_attention_mask: torch.Tensor, | |
| valid_pair_mask: torch.Tensor | None = None, | |
| ) -> PRISMOutput: | |
| """ | |
| pixel_values: ``(B, T, 3, H, W)`` — padded to T in the collator. | |
| valid_mask: ``(B, T)`` bool, True = real frame. | |
| composed_input_ids: ``(B, L)`` — tokenized recomposed caption per pair. | |
| composed_attention_mask: ``(B, L)``. | |
| valid_pair_mask: ``(B,)`` bool — True if the composer succeeded. | |
| """ | |
| B = pixel_values.shape[0] | |
| device = pixel_values.device | |
| use_ema = getattr(self.config, "use_ema", False) | |
| # ---- Frozen encoders ---- | |
| patches = self._encode_video(pixel_values) # (B, T, P, d_v) | |
| T = patches.shape[1] | |
| e_text = self._encode_text(composed_input_ids, composed_attention_mask) # (B, d_t) | |
| # ---- Decompositional encoder θ (and EMA θ̄ for temporal targets) ---- | |
| kpm = ~valid_mask | |
| z_vi_seq, z_vv_seq = self.encoder(patches, key_padding_mask=kpm) | |
| if use_ema: | |
| with torch.no_grad(): | |
| z_vi_seq_ema, z_vv_seq_ema = self.target_encoder(patches, key_padding_mask=kpm) | |
| # ---- Cross-pairing: view-variant from clip i, view-invariant from (i+1) ---- | |
| shift = torch.roll(torch.arange(B, device=device), shifts=-1, dims=0) | |
| z_vv = z_vv_seq # view-variant (clip i) | |
| z_vi = z_vi_seq[shift] # view-invariant (clip i+1) | |
| valid_vv = valid_mask | |
| valid_vi = valid_mask[shift] | |
| if use_ema: | |
| z_vv_ema = z_vv_seq_ema | |
| z_vi_ema = z_vi_seq_ema[shift] | |
| # ---- Sliding-shift augmentation (training only) ---- | |
| if self.training and getattr(self.config, "sliding_shift_aug", True): | |
| plan = _sample_shift_plan(valid_vv, valid_vi, T) | |
| z_vv = _apply_shift(z_vv, plan["src_idx_a"], plan["new_valid_a"]) | |
| z_vi = _apply_shift(z_vi, plan["src_idx_b"], plan["new_valid_b"]) | |
| if use_ema: | |
| z_vv_ema = _apply_shift(z_vv_ema, plan["src_idx_a"], plan["new_valid_a"]) | |
| z_vi_ema = _apply_shift(z_vi_ema, plan["src_idx_b"], plan["new_valid_b"]) | |
| valid_vv = plan["new_valid_a"] | |
| valid_vi = plan["new_valid_b"] | |
| # ---- Compositional predictor φ ---- | |
| out = self.predictor(z_vv, z_vi, valid_vv=valid_vv, valid_vi=valid_vi) | |
| s = out["s"] # (B, d_t) compositional latent | |
| z_vi_pred = out["z_vi_pred"] # (B, T, d_z) vi next-frame head | |
| z_vv_pred = out["z_vv_pred"] # (B, T, d_z) vv next-frame head | |
| pair_valid = out["pair_valid"] # (B, T) | |
| # ---- L_decomp: InfoNCE(s, e_text) over valid pairs ---- | |
| if valid_pair_mask is None: | |
| valid_pair_mask = torch.ones(B, dtype=torch.bool, device=device) | |
| valid_pair_mask = valid_pair_mask & pair_valid.any(dim=1) | |
| if self.config.infonce_all_gather: | |
| s_g = _all_gather_with_grad(s) | |
| e_g = _all_gather_with_grad(e_text) | |
| valid_g = _all_gather_bool(valid_pair_mask) | |
| else: | |
| s_g, e_g, valid_g = s, e_text, valid_pair_mask | |
| valid_idx = valid_g.nonzero(as_tuple=True)[0] | |
| n_valid = int(valid_idx.numel()) | |
| if n_valid >= 2: | |
| loss_decomp = _symmetric_infonce(s_g[valid_idx], e_g[valid_idx], self.logit_scale) | |
| else: | |
| loss_decomp = torch.zeros((), device=device) | |
| # ---- L_temp: 1 - cos(prediction at t, EMA target at t+1), per stream ---- | |
| if T >= 2: | |
| if use_ema: | |
| tgt_vi = z_vi_ema[:, 1:T, :] | |
| tgt_vv = z_vv_ema[:, 1:T, :] | |
| else: | |
| tgt_vi = z_vi.detach()[:, 1:T, :] | |
| tgt_vv = z_vv.detach()[:, 1:T, :] | |
| valid_next_vi = valid_vi[:, :T - 1] & valid_vi[:, 1:T] | |
| valid_next_vv = valid_vv[:, :T - 1] & valid_vv[:, 1:T] | |
| err_vi = 1.0 - F.cosine_similarity(z_vi_pred[:, :T - 1, :], tgt_vi, dim=-1) | |
| err_vv = 1.0 - F.cosine_similarity(z_vv_pred[:, :T - 1, :], tgt_vv, dim=-1) | |
| loss_temp_vi = err_vi[valid_next_vi].mean() if valid_next_vi.any() else torch.zeros((), device=device) | |
| loss_temp_vv = err_vv[valid_next_vv].mean() if valid_next_vv.any() else torch.zeros((), device=device) | |
| else: | |
| loss_temp_vi = torch.zeros((), device=device) | |
| loss_temp_vv = torch.zeros((), device=device) | |
| loss_temp = 0.5 * (loss_temp_vi + loss_temp_vv) | |
| loss = self.config.lambda_decomp * loss_decomp + self.config.lambda_temp * loss_temp | |
| return PRISMOutput( | |
| loss=loss, | |
| loss_decomp=loss_decomp.detach(), | |
| loss_temp_vi=loss_temp_vi.detach(), | |
| loss_temp_vv=loss_temp_vv.detach(), | |
| n_valid_pairs=n_valid, | |
| ) | |