Integrate with Sentence Transformers via MultiVectorEncoder, fix transformers 5.x processor and modeling

#5
by tomaarsen HF Staff - opened
1_MultiVectorMask/config.json ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ {
2
+ "skiplist_words": [],
3
+ "skiplist_tasks": [],
4
+ "keep_only_token_ids": null
5
+ }
README.md CHANGED
@@ -6,6 +6,8 @@ pipeline_tag: visual-document-retrieval
6
  library_name: transformers
7
  tags:
8
  - transformers
 
 
9
  - multimodal_embedding
10
  - embedding
11
  - colpali
@@ -30,6 +32,47 @@ The model is trained using a multi-stage strategy that combines large-scale text
30
 
31
  ## Usage
32
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
33
  **Requirements**
34
  ```
35
  pillow
 
6
  library_name: transformers
7
  tags:
8
  - transformers
9
+ - sentence-transformers
10
+ - multi-vector
11
  - multimodal_embedding
12
  - embedding
13
  - colpali
 
32
 
33
  ## Usage
34
 
35
+ ### Sentence Transformers
36
+
37
+ This model can be used with [Sentence Transformers](https://www.sbert.net/) as a multi-vector (ColBERT-style late interaction) retriever via the `MultiVectorEncoder`:
38
+
39
+ ```bash
40
+ pip install "sentence-transformers[image]>=6.0.0"
41
+ ```
42
+
43
+ ```python
44
+ from sentence_transformers import MultiVectorEncoder
45
+
46
+ model = MultiVectorEncoder(
47
+ "OpenSearch-AI/Ops-Colqwen3-4B",
48
+ trust_remote_code=True,
49
+ model_kwargs={"dtype": "bfloat16"},
50
+ )
51
+
52
+ queries = [
53
+ "What is the variable represented on the y-axis of the graph?",
54
+ "Total outlay is maximum in which year?",
55
+ ]
56
+ images = [
57
+ "https://huggingface.co/datasets/sentence-transformers/example-documents/resolve/main/doc1.jpg",
58
+ "https://huggingface.co/datasets/sentence-transformers/example-documents/resolve/main/doc2.jpg",
59
+ ]
60
+
61
+ query_embeddings = model.encode_query(queries)
62
+ image_embeddings = model.encode_document(images)
63
+ print(query_embeddings[0].shape, image_embeddings[0].shape)
64
+ # torch.Size([25, 2560]) torch.Size([1254, 2560])
65
+
66
+ # Diagonal should have higher scores
67
+ scores = model.similarity(query_embeddings, image_embeddings)
68
+ print(scores)
69
+ # tensor([[17.4668, 12.9785],
70
+ # [ 7.0088, 15.9492]], device='cuda:0')
71
+ ```
72
+
73
+ ### Transformers
74
+
75
+
76
  **Requirements**
77
  ```
78
  pillow
config_sentence_transformers.json ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "model_type": "MultiVectorEncoder",
3
+ "similarity_fn_name": "maxsim",
4
+ "prompts": {},
5
+ "default_prompt_name": null,
6
+ "__version__": {
7
+ "sentence_transformers": "6.0.0"
8
+ }
9
+ }
modeling_ops_colqwen3.py CHANGED
@@ -54,7 +54,7 @@ class OpsColQwen3Model(OpsColQwen3PreTrainedModel):
54
  model.dims = dims
55
  return model
56
 
57
- def forward(self, input_ids: Optional[torch.LongTensor] = None, attention_mask: Optional[torch.Tensor] = None, pixel_values: Optional[torch.Tensor] = None, image_grid_thw: Optional[torch.Tensor] = None, **kwargs) -> torch.Tensor:
58
  has_pixel_values = pixel_values is not None
59
 
60
  if has_pixel_values:
@@ -67,6 +67,11 @@ class OpsColQwen3Model(OpsColQwen3PreTrainedModel):
67
  unpadded = [pixel_sequence[: int(offset.item())] for pixel_sequence, offset in zip(pixel_values, offsets)]
68
  pixel_values = torch.cat(unpadded, dim=0) if unpadded else None
69
 
 
 
 
 
 
70
  outputs = self.qwen3vl(
71
  input_ids=input_ids,
72
  attention_mask=attention_mask,
@@ -75,6 +80,7 @@ class OpsColQwen3Model(OpsColQwen3PreTrainedModel):
75
  use_cache=False,
76
  output_hidden_states=True,
77
  return_dict=True,
 
78
  )
79
 
80
  last_hidden_states = outputs.last_hidden_state
 
54
  model.dims = dims
55
  return model
56
 
57
+ def forward(self, input_ids: Optional[torch.LongTensor] = None, attention_mask: Optional[torch.Tensor] = None, pixel_values: Optional[torch.Tensor] = None, image_grid_thw: Optional[torch.Tensor] = None, mm_token_type_ids: Optional[torch.Tensor] = None, **kwargs) -> torch.Tensor:
58
  has_pixel_values = pixel_values is not None
59
 
60
  if has_pixel_values:
 
67
  unpadded = [pixel_sequence[: int(offset.item())] for pixel_sequence, offset in zip(pixel_values, offsets)]
68
  pixel_values = torch.cat(unpadded, dim=0) if unpadded else None
69
 
70
+ extra_kwargs = {}
71
+ if mm_token_type_ids is not None:
72
+ # Required for M-RoPE on recent transformers versions. Passed conditionally so
73
+ # older versions without the argument keep working.
74
+ extra_kwargs["mm_token_type_ids"] = mm_token_type_ids
75
  outputs = self.qwen3vl(
76
  input_ids=input_ids,
77
  attention_mask=attention_mask,
 
80
  use_cache=False,
81
  output_hidden_states=True,
82
  return_dict=True,
83
+ **extra_kwargs,
84
  )
85
 
86
  last_hidden_states = outputs.last_hidden_state
modules.json ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [
2
+ {
3
+ "idx": 0,
4
+ "name": "0",
5
+ "path": "",
6
+ "type": "sentence_transformers.base.modules.transformer.Transformer"
7
+ },
8
+ {
9
+ "idx": 1,
10
+ "name": "1",
11
+ "path": "1_MultiVectorMask",
12
+ "type": "sentence_transformers.multi_vector_encoder.modules.multi_vector_mask.MultiVectorMask"
13
+ }
14
+ ]
processing_ops_colqwen3.py CHANGED
@@ -61,13 +61,44 @@ class OpsColQwen3Processor(Qwen3VLProcessor):
61
  if self.tokenizer is not None:
62
  self.tokenizer.padding_side = "left"
63
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
64
  def process_images(self, images: List[Image.Image], return_tensors: str = "pt", **kwargs) -> Union[BatchFeature, BatchEncoding]:
65
  """
66
  Process a batch of PIL images for the model.
67
  """
68
  images = [image.convert("RGB") for image in images]
69
 
70
- batch_doc = self(text=[self.visual_prompt_prefix] * len(images), images=images, padding="longest", return_tensors=return_tensors, **kwargs)
71
 
72
  if batch_doc["pixel_values"].numel() == 0:
73
  return batch_doc
@@ -83,7 +114,7 @@ class OpsColQwen3Processor(Qwen3VLProcessor):
83
  Process a list of text queries.
84
  """
85
  processed_queries = [self.query_prefix + q + self.query_augmentation_token * 10 for q in queries]
86
- return self(text=processed_queries, return_tensors=return_tensors, padding="longest", **kwargs)
87
 
88
  @staticmethod
89
  def score_multi_vector(
 
61
  if self.tokenizer is not None:
62
  self.tokenizer.padding_side = "left"
63
 
64
+ def __call__(self, images=None, text=None, audio=None, videos=None, **kwargs) -> Union[BatchFeature, BatchEncoding]:
65
+ """Standard processor interface with retrieval formatting.
66
+
67
+ Routes plain calls through process_images/process_queries so that generic pipelines
68
+ (for example Sentence Transformers) produce the same inputs as the dedicated methods:
69
+ images get the visual prompt and per-image pixel padding, text gets the query prefix
70
+ and augmentation tokens. `mm_token_type_ids` is requested so the Qwen3-VL M-RoPE
71
+ requirement of recent transformers versions is satisfied.
72
+ """
73
+ # The process_* methods set padding and return_tensors themselves: drop duplicates that
74
+ # generic callers pass through nested kwargs.
75
+ for nest in ("text_kwargs", "images_kwargs", "videos_kwargs", "audio_kwargs", "common_kwargs"):
76
+ sub = kwargs.get(nest)
77
+ if isinstance(sub, dict):
78
+ for key in ("padding", "return_tensors", "return_mm_token_type_ids"):
79
+ sub.pop(key, None)
80
+ if images is not None:
81
+ image_list = images if isinstance(images, list) else [images]
82
+ flat_images = []
83
+ for item in image_list:
84
+ flat_images.extend(item if isinstance(item, list) else [item])
85
+ return self.process_images(flat_images, return_mm_token_type_ids=True, **kwargs)
86
+ if text is not None:
87
+ texts = text if isinstance(text, str) else list(text)
88
+ return self.process_queries([texts] if isinstance(texts, str) else texts, **kwargs)
89
+ raise ValueError("You have to specify at least one of `images` or `text`.")
90
+
91
+ def _raw_call(self, *args, **kwargs) -> Union[BatchFeature, BatchEncoding]:
92
+ """The inherited Qwen3VLProcessor call, used internally by the process_* methods."""
93
+ return super().__call__(*args, **kwargs)
94
+
95
  def process_images(self, images: List[Image.Image], return_tensors: str = "pt", **kwargs) -> Union[BatchFeature, BatchEncoding]:
96
  """
97
  Process a batch of PIL images for the model.
98
  """
99
  images = [image.convert("RGB") for image in images]
100
 
101
+ batch_doc = self._raw_call(text=[self.visual_prompt_prefix] * len(images), images=images, padding="longest", return_tensors=return_tensors, **kwargs)
102
 
103
  if batch_doc["pixel_values"].numel() == 0:
104
  return batch_doc
 
114
  Process a list of text queries.
115
  """
116
  processed_queries = [self.query_prefix + q + self.query_augmentation_token * 10 for q in queries]
117
+ return self._raw_call(text=processed_queries, return_tensors=return_tensors, padding="longest", **kwargs)
118
 
119
  @staticmethod
120
  def score_multi_vector(
sentence_bert_config.json ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "transformer_task": "retrieval",
3
+ "modality_config": {
4
+ "text": {
5
+ "method": "forward",
6
+ "method_output_name": null
7
+ },
8
+ "image": {
9
+ "method": "forward",
10
+ "method_output_name": null
11
+ }
12
+ },
13
+ "module_output_name": "token_embeddings"
14
+ }