Feature Extraction
sentence-transformers
Safetensors
English
modernvbert
sparse-retrieval
splade
visual-document-retrieval
multimodal
information-retrieval
inference-free
sparse-encoder
custom_code
Instructions to use naver/v-splade-efficient with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- sentence-transformers
How to use naver/v-splade-efficient with sentence-transformers:
from sentence_transformers import SentenceTransformer model = SentenceTransformer("naver/v-splade-efficient", trust_remote_code=True) sentences = [ "The weather is lovely today.", "It's so sunny outside!", "He drove to the stadium." ] embeddings = model.encode(sentences) similarities = model.similarity(embeddings, embeddings) print(similarities.shape) # [3, 3] - Notebooks
- Google Colab
- Kaggle
File size: 7,541 Bytes
b7875cd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 | """V-SPLADE document encoder for Hugging Face Transformers.
Wraps the ModernVBERT backbone (``transformers>=5.3.0``) together with the
V-SPLADE MLM head so that this repository loads directly with
``AutoModelForMaskedLM.from_pretrained(..., trust_remote_code=True)``.
The module tree deliberately mirrors the checkpoint layout of the V-SPLADE
export (``encoder.encoder.model.*`` for the backbone, ``encoder.mlm_head.*``
for the sparse head), so ``model.safetensors`` loads without any key
remapping. The ``query_encoder.*`` tensors hold the inference-free Li-LSR
query lookup (used by the Sentence Transformers integration) and are not part
of the document encoder, so they are ignored here.
The returned ``logits`` are the SPLADE term logits: MLM logits scaled by
``hidden_size ** -0.25`` with special tokens masked out, exactly as in
https://github.com/naver/v-splade (``UnifiedRetriever._apply_sparse_head``).
A sparse document embedding is obtained via ``log1p(relu(logits))`` followed
by a max-pool over the sequence dimension (see the README).
"""
from __future__ import annotations
import torch
import torch.nn.functional as F
from torch import nn
from transformers.modeling_outputs import MaskedLMOutput
try:
from transformers.models.modernvbert.configuration_modernvbert import ModernVBertConfig
from transformers.models.modernvbert.modeling_modernvbert import (
ModernVBertModel,
ModernVBertPreTrainedModel,
)
except ImportError as exc:
raise ImportError(
"V-SPLADE requires the ModernVBERT architecture, which is available in "
"transformers>=5.3.0. Please upgrade with `pip install -U transformers`."
) from exc
# Special tokens that are masked out of the sparse representation:
# [UNK], [CLS], [SEP], [PAD], [MASK]
SPECIAL_TOKEN_IDS = [50280, 50281, 50282, 50283, 50284]
class VSPLADEDecoupledEmbedding(nn.Embedding):
"""Word embeddings split into the base vocabulary and the added vision tokens.
Matches the V-SPLADE export layout: ``weight`` holds the base (MLM) vocabulary
and ``additional_embedding.weight`` holds the extra tokens appended for the
vision chat format (``<image>``, ``<end_of_utterance>``, tile markers, ...).
"""
def __init__(self, num_embeddings: int, num_additional_embeddings: int, embedding_dim: int, **kwargs) -> None:
super().__init__(num_embeddings, embedding_dim, **kwargs)
self.num_additional_embeddings = num_additional_embeddings
self.additional_embedding = nn.Embedding(num_additional_embeddings, embedding_dim)
def forward(self, input_ids: torch.LongTensor) -> torch.Tensor:
input_ids = input_ids.clone()
additional_vocab_indices = torch.where(input_ids >= self.num_embeddings)
additional_embeddings = self.additional_embedding(input_ids[additional_vocab_indices] - self.num_embeddings)
input_ids[additional_vocab_indices] = 0
full_vector = F.embedding(input_ids, self.weight)
full_vector[additional_vocab_indices] = additional_embeddings
return full_vector
class VSPLADEModalityProjection(nn.Module):
"""Vision-to-text projection stored as ``modality_projection.proj`` in the export."""
def __init__(self, input_size: int, output_size: int) -> None:
super().__init__()
self.proj = nn.Linear(input_size, output_size, bias=False)
@property
def weight(self) -> torch.Tensor:
# ModernVBertPreTrainedModel._init_weights initializes ``modality_projection.weight``
return self.proj.weight
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
return self.proj(hidden_states)
class VSPLADEMLMHead(nn.Module):
"""V-SPLADE MLM head: dense -> GELU -> LayerNorm -> decoder (base vocabulary)."""
def __init__(self, hidden_size: int, vocab_size: int) -> None:
super().__init__()
self.dense = nn.Linear(hidden_size, hidden_size)
self.norm = nn.LayerNorm(hidden_size)
self.decoder = nn.Linear(hidden_size, vocab_size)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
return self.decoder(self.norm(F.gelu(self.dense(hidden_states))))
class _Wrapper(nn.Module):
"""Empty container used to mirror the checkpoint's key prefixes."""
class VSPLADEForMaskedLM(ModernVBertPreTrainedModel):
config_class = ModernVBertConfig
_keys_to_ignore_on_load_unexpected = [r"query_encoder\..*"]
def __init__(self, config: ModernVBertConfig) -> None:
super().__init__(config)
main_vocab_size = config.text_config.vocab_size - config.additional_vocab_size
backbone = ModernVBertModel(config)
# The export stores the connector projection under an extra ``proj`` level; mirror that.
backbone.connector.modality_projection = VSPLADEModalityProjection(
input_size=config.vision_config.hidden_size * (config.pixel_shuffle_factor**2),
output_size=config.text_config.hidden_size,
)
# The export splits the embedding into base + additional tokens; mirror that.
backbone.text_model.set_input_embeddings(
VSPLADEDecoupledEmbedding(
num_embeddings=main_vocab_size,
num_additional_embeddings=config.additional_vocab_size,
embedding_dim=config.text_config.hidden_size,
padding_idx=getattr(config, "pad_token_id", None),
)
)
self.encoder = _Wrapper()
self.encoder.encoder = _Wrapper()
self.encoder.encoder.model = backbone
self.encoder.mlm_head = VSPLADEMLMHead(config.text_config.hidden_size, main_vocab_size)
self.logit_scale = config.text_config.hidden_size**-0.25
self.post_init()
def get_input_embeddings(self):
return self.encoder.encoder.model.get_input_embeddings()
def set_input_embeddings(self, value):
self.encoder.encoder.model.set_input_embeddings(value)
def forward(
self,
input_ids: torch.LongTensor | None = None,
attention_mask: torch.Tensor | None = None,
position_ids: torch.LongTensor | None = None,
inputs_embeds: torch.FloatTensor | None = None,
pixel_values: torch.FloatTensor | None = None,
pixel_attention_mask: torch.BoolTensor | None = None,
image_hidden_states: torch.FloatTensor | None = None,
return_dict: bool | None = None,
) -> MaskedLMOutput:
outputs = self.encoder.encoder.model(
input_ids=input_ids,
attention_mask=attention_mask,
position_ids=position_ids,
inputs_embeds=inputs_embeds,
pixel_values=pixel_values,
pixel_attention_mask=pixel_attention_mask,
image_hidden_states=image_hidden_states,
return_dict=True,
)
logits = self.encoder.mlm_head(outputs.last_hidden_state) * self.logit_scale
# Zero out special tokens so they never activate in the sparse representation
# (log1p(relu(0)) == 0), matching the reference special_token_mask.
# Built on the fly: buffers created in __init__ do not survive meta-device loading.
special_token_ids = torch.tensor(SPECIAL_TOKEN_IDS, dtype=torch.long, device=logits.device)
logits = logits.index_fill(-1, special_token_ids, 0.0)
return MaskedLMOutput(
logits=logits,
hidden_states=outputs.hidden_states,
attentions=outputs.attentions,
)
__all__ = ["VSPLADEForMaskedLM"]
|