transformers-haystack 0.2.0__tar.gz → 1.0.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (37) hide show
  1. transformers_haystack-1.0.0/CHANGELOG.md +40 -0
  2. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/PKG-INFO +4 -4
  3. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/pyproject.toml +4 -2
  4. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/common/transformers/utils.py +19 -20
  5. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/components/classifiers/transformers/zero_shot_document_classifier.py +9 -11
  6. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/components/extractors/transformers/named_entity_extractor.py +10 -27
  7. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/components/generators/transformers/chat/chat_generator.py +50 -65
  8. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/components/readers/transformers/extractive_reader.py +24 -18
  9. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/components/routers/transformers/text_router.py +13 -15
  10. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/components/routers/transformers/zero_shot_text_router.py +14 -12
  11. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/tests/test_chat_generator.py +251 -172
  12. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/tests/test_extractive_reader.py +31 -5
  13. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/tests/test_named_entity_extractor.py +23 -5
  14. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/tests/test_text_router.py +41 -23
  15. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/tests/test_utils.py +26 -0
  16. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/tests/test_zero_shot_document_classifier.py +43 -21
  17. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/tests/test_zero_shot_text_router.py +34 -5
  18. transformers_haystack-0.2.0/CHANGELOG.md +0 -9
  19. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/.gitignore +0 -0
  20. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/LICENSE.txt +0 -0
  21. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/README.md +0 -0
  22. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/pydoc/config_docusaurus.yml +0 -0
  23. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/common/py.typed +0 -0
  24. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/common/transformers/__init__.py +0 -0
  25. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/components/classifiers/py.typed +0 -0
  26. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/components/classifiers/transformers/__init__.py +0 -0
  27. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/components/extractors/py.typed +0 -0
  28. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/components/extractors/transformers/__init__.py +0 -0
  29. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/components/generators/py.typed +0 -0
  30. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/components/generators/transformers/__init__.py +0 -0
  31. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/components/generators/transformers/chat/__init__.py +0 -0
  32. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/components/readers/py.typed +0 -0
  33. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/components/readers/transformers/__init__.py +0 -0
  34. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/components/routers/py.typed +0 -0
  35. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/src/haystack_integrations/components/routers/transformers/__init__.py +0 -0
  36. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/tests/__init__.py +0 -0
  37. {transformers_haystack-0.2.0 → transformers_haystack-1.0.0}/tests/conftest.py +0 -0
@@ -0,0 +1,40 @@
1
+ # Changelog
2
+
3
+ ## [integrations/transformers-v0.3.0] - 2026-08-24
4
+
5
+ ### 🐛 Bug Fixes
6
+
7
+ - Fix new issues raised by ruff 0.16.0 (#3670)
8
+ - Serialize missing init params in to_dict (transformers, huggingface_api, ragas) (#3808)
9
+
10
+ ### 🧹 Chores
11
+
12
+ - Clarify how generation_kwargs passed in run are handled (#3805)
13
+
14
+
15
+ ## [integrations/transformers-v0.2.0] - 2026-07-06
16
+
17
+ ### 📚 Documentation
18
+
19
+ - Replace old haystack core imports with haystack_integrations paths (#3545)
20
+
21
+ ### 🧪 Testing
22
+
23
+ - Improve del_hf_env_vars fixture (#3428)
24
+ - Trust test modules under Haystack 3.0's deserialization allowlist (#3537)
25
+ - Make Tool/Agent serialization assertions version-agnostic for Haystack 2.x/3.x (#3533)
26
+ - Force Transformers and Sentence Transformers integration tests to run on CPU (#3550)
27
+
28
+ ### 🧹 Chores
29
+
30
+ - Improve consistency of integrations folder structure (#3430)
31
+ - Support sync streaming callbacks in async contexts for Haystack 2.x/3.x compatibility (#3534)
32
+
33
+
34
+ ## [integrations/transformers-v0.1.0] - 2026-06-08
35
+
36
+ ### 🚀 Features
37
+
38
+ - Move Transformers components from Haystack (#3409)
39
+
40
+ <!-- generated by git-cliff -->
@@ -1,6 +1,6 @@
1
- Metadata-Version: 2.4
1
+ Metadata-Version: 2.5
2
2
  Name: transformers-haystack
3
- Version: 0.2.0
3
+ Version: 1.0.0
4
4
  Summary: Haystack integration for transformers
5
5
  Project-URL: Documentation, https://github.com/deepset-ai/haystack-core-integrations/tree/main/integrations/transformers#readme
6
6
  Project-URL: Issues, https://github.com/deepset-ai/haystack-core-integrations/issues
@@ -19,8 +19,8 @@ Classifier: Programming Language :: Python :: 3.14
19
19
  Classifier: Programming Language :: Python :: Implementation :: CPython
20
20
  Classifier: Programming Language :: Python :: Implementation :: PyPy
21
21
  Requires-Python: >=3.10
22
- Requires-Dist: haystack-ai>=2.30.0
23
- Requires-Dist: transformers[sentencepiece,torch]>=4.57.0
22
+ Requires-Dist: haystack-ai>=3.0.0
23
+ Requires-Dist: transformers[sentencepiece,torch]>=5.0.0
24
24
  Description-Content-Type: text/markdown
25
25
 
26
26
  # transformers-haystack
@@ -23,7 +23,7 @@ classifiers = [
23
23
  "Programming Language :: Python :: Implementation :: CPython",
24
24
  "Programming Language :: Python :: Implementation :: PyPy",
25
25
  ]
26
- dependencies = ["haystack-ai>=2.30.0", "transformers[torch,sentencepiece]>=4.57.0"]
26
+ dependencies = ["haystack-ai>=3.0.0", "transformers[torch,sentencepiece]>=5.0.0"]
27
27
 
28
28
  [project.urls]
29
29
  Documentation = "https://github.com/deepset-ai/haystack-core-integrations/tree/main/integrations/transformers#readme"
@@ -81,8 +81,9 @@ disallow_incomplete_defs = true
81
81
 
82
82
  [[tool.mypy.overrides]]
83
83
  module = [
84
- # these libraries do not ship type stubs
84
+ # accelerate does not ship type stubs
85
85
  "accelerate.*",
86
+ # tokenizers only ships a py.typed marker from 0.23.1 on; kept for older versions
86
87
  "tokenizers.*",
87
88
  ]
88
89
  ignore_missing_imports = true
@@ -140,6 +141,7 @@ ignore = [
140
141
  "PLR0912",
141
142
  "PLR0913",
142
143
  "PLR0915",
144
+ "PLR0917",
143
145
  # Allow `Any` type - used legitimately for dynamic types and SDK boundaries
144
146
  "ANN401",
145
147
  ]
@@ -3,7 +3,6 @@
3
3
  # SPDX-License-Identifier: Apache-2.0
4
4
 
5
5
  import asyncio
6
- import copy
7
6
  import inspect
8
7
  from typing import Any
9
8
 
@@ -12,7 +11,6 @@ from haystack import logging
12
11
  from haystack.dataclasses import ComponentInfo, StreamingCallbackT, StreamingChunk, SyncStreamingCallbackT
13
12
  from haystack.utils.auth import Secret
14
13
  from haystack.utils.device import ComponentDevice
15
- from huggingface_hub import model_info
16
14
 
17
15
  from transformers import (
18
16
  PreTrainedTokenizer,
@@ -25,6 +23,21 @@ from transformers import (
25
23
  logger = logging.getLogger(__name__)
26
24
 
27
25
 
26
+ def _with_hf_token(hf_kwargs: dict[str, Any], token: Secret | None) -> dict[str, Any]:
27
+ """
28
+ Return a copy of Hugging Face keyword arguments with a resolved token.
29
+
30
+ An explicitly provided `token` in `hf_kwargs` takes precedence over the `Secret`.
31
+
32
+ :param hf_kwargs: Keyword arguments passed to a Hugging Face API.
33
+ :param token: The token to resolve when `hf_kwargs` does not already contain one.
34
+ """
35
+ resolved_kwargs = hf_kwargs.copy()
36
+ if "token" not in resolved_kwargs:
37
+ resolved_kwargs["token"] = token.resolve_value() if token else None
38
+ return resolved_kwargs
39
+
40
+
28
41
  def _resolve_hf_device_map(device: ComponentDevice | None, model_kwargs: dict[str, Any] | None) -> dict[str, Any]:
29
42
  """
30
43
  Update `model_kwargs` to include the keyword argument `device_map`.
@@ -40,7 +53,7 @@ def _resolve_hf_device_map(device: ComponentDevice | None, model_kwargs: dict[st
40
53
  :param model_kwargs: Additional HF keyword arguments passed to `AutoModel.from_pretrained`.
41
54
  For details on what kwargs you can pass, see the model's documentation.
42
55
  """
43
- model_kwargs = copy.copy(model_kwargs) or {}
56
+ model_kwargs = dict(model_kwargs or {})
44
57
  if model_kwargs.get("device_map"):
45
58
  if device is not None:
46
59
  logger.warning(
@@ -62,10 +75,8 @@ def _resolve_hf_device_map(device: ComponentDevice | None, model_kwargs: dict[st
62
75
  def _resolve_hf_pipeline_kwargs(
63
76
  huggingface_pipeline_kwargs: dict[str, Any],
64
77
  model: str,
65
- task: str | None,
66
- supported_tasks: list[str],
78
+ task: str,
67
79
  device: ComponentDevice | None,
68
- token: Secret | None,
69
80
  ) -> dict[str, Any]:
70
81
  """
71
82
  Resolve the HuggingFace pipeline keyword arguments based on explicit user inputs.
@@ -74,30 +85,18 @@ def _resolve_hf_pipeline_kwargs(
74
85
  Hugging Face pipeline.
75
86
  :param model: The name or path of a Hugging Face model for on the HuggingFace Hub.
76
87
  :param task: The task for the Hugging Face pipeline.
77
- :param supported_tasks: The list of supported tasks to check the task of the model against. If the task of the model
78
- is not present within this list then a ValueError is thrown.
79
88
  :param device: The device on which the model is loaded. If `None`, the default device is automatically
80
89
  selected. If a device/device map is specified in `huggingface_pipeline_kwargs`, it overrides this parameter.
81
- :param token: The token to use as HTTP bearer authorization for remote files.
82
- If the token is also specified in the `huggingface_pipeline_kwargs`, this parameter will be ignored.
83
90
  """
84
- resolved_token = token.resolve_value() if token else None
91
+ huggingface_pipeline_kwargs = huggingface_pipeline_kwargs.copy()
92
+
85
93
  # check if the huggingface_pipeline_kwargs contain the essential parameters
86
94
  # otherwise, populate them with values from other init parameters
87
95
  huggingface_pipeline_kwargs.setdefault("model", model)
88
- huggingface_pipeline_kwargs.setdefault("token", resolved_token)
89
96
 
90
97
  resolved_device = ComponentDevice.resolve_device(device)
91
98
  resolved_device.update_hf_kwargs(huggingface_pipeline_kwargs, overwrite=False)
92
99
 
93
- # task identification and validation
94
- task = task or huggingface_pipeline_kwargs.get("task")
95
- if task is None and isinstance(huggingface_pipeline_kwargs["model"], str):
96
- task = model_info(huggingface_pipeline_kwargs["model"], token=huggingface_pipeline_kwargs["token"]).pipeline_tag
97
-
98
- if task not in supported_tasks:
99
- msg = f"Task '{task}' is not supported. The supported tasks are: {', '.join(supported_tasks)}."
100
- raise ValueError(msg)
101
100
  huggingface_pipeline_kwargs["task"] = task
102
101
  return huggingface_pipeline_kwargs
103
102
 
@@ -9,7 +9,7 @@ from haystack import Document, component, default_from_dict, default_to_dict
9
9
  from haystack.utils import ComponentDevice, Secret
10
10
  from haystack.utils.hf import deserialize_hf_model_kwargs, serialize_hf_model_kwargs
11
11
 
12
- from haystack_integrations.common.transformers.utils import _resolve_hf_pipeline_kwargs
12
+ from haystack_integrations.common.transformers.utils import _resolve_hf_pipeline_kwargs, _with_hf_token
13
13
  from transformers import Pipeline as HfPipeline
14
14
  from transformers import pipeline
15
15
 
@@ -123,9 +123,7 @@ class TransformersZeroShotDocumentClassifier:
123
123
  huggingface_pipeline_kwargs=huggingface_pipeline_kwargs or {},
124
124
  model=model,
125
125
  task="zero-shot-classification",
126
- supported_tasks=["zero-shot-classification"],
127
126
  device=device,
128
- token=token,
129
127
  )
130
128
 
131
129
  self.huggingface_pipeline_kwargs = huggingface_pipeline_kwargs
@@ -143,8 +141,11 @@ class TransformersZeroShotDocumentClassifier:
143
141
  """
144
142
  Initializes the component.
145
143
  """
146
- if self.pipeline is None:
147
- self.pipeline = pipeline(**self.huggingface_pipeline_kwargs)
144
+ if self.pipeline is not None:
145
+ return
146
+
147
+ pipeline_kwargs = _with_hf_token(self.huggingface_pipeline_kwargs, self.token)
148
+ self.pipeline = pipeline(**pipeline_kwargs)
148
149
 
149
150
  def to_dict(self) -> dict[str, Any]:
150
151
  """
@@ -201,8 +202,8 @@ class TransformersZeroShotDocumentClassifier:
201
202
  - `documents`: A list of documents with an added metadata field called `classification`.
202
203
  """
203
204
 
204
- if self.pipeline is None:
205
- self.warm_up()
205
+ self.warm_up()
206
+ assert self.pipeline is not None # noqa: S101
206
207
 
207
208
  if not isinstance(documents, list) or (documents and not isinstance(documents[0], Document)):
208
209
  msg = (
@@ -229,10 +230,7 @@ class TransformersZeroShotDocumentClassifier:
229
230
  for doc in documents
230
231
  ]
231
232
 
232
- # mypy doesn't know this is set in warm_up
233
- predictions = self.pipeline( # type: ignore[misc]
234
- texts, self.labels, multi_label=self.multi_label, batch_size=batch_size
235
- )
233
+ predictions = self.pipeline(texts, self.labels, multi_label=self.multi_label, batch_size=batch_size)
236
234
 
237
235
  new_documents = []
238
236
  for prediction, document in zip(predictions, documents, strict=True):
@@ -10,9 +10,9 @@ from haystack.utils.auth import Secret
10
10
  from haystack.utils.device import ComponentDevice
11
11
  from haystack.utils.hf import deserialize_hf_model_kwargs, serialize_hf_model_kwargs
12
12
 
13
- from haystack_integrations.common.transformers.utils import _resolve_hf_pipeline_kwargs
14
- from transformers import AutoModelForTokenClassification, AutoTokenizer, pipeline
13
+ from haystack_integrations.common.transformers.utils import _resolve_hf_pipeline_kwargs, _with_hf_token
15
14
  from transformers import Pipeline as HfPipeline
15
+ from transformers import pipeline
16
16
 
17
17
 
18
18
  @dataclass
@@ -97,15 +97,10 @@ class TransformersNamedEntityExtractor:
97
97
  huggingface_pipeline_kwargs=pipeline_kwargs or {},
98
98
  model=model,
99
99
  task="ner",
100
- supported_tasks=["ner"],
101
100
  device=self.device,
102
- token=token,
103
101
  )
104
102
 
105
- self.tokenizer: Any = None
106
- self.model: AutoModelForTokenClassification | None = None
107
103
  self.pipeline: HfPipeline | None = None
108
- self._warmed_up: bool = False
109
104
 
110
105
  def warm_up(self) -> None:
111
106
  """
@@ -114,24 +109,13 @@ class TransformersNamedEntityExtractor:
114
109
  :raises ComponentError:
115
110
  If the component fails to initialize successfully.
116
111
  """
117
- if self._warmed_up:
112
+ if self.pipeline is not None:
118
113
  return
119
114
 
120
115
  try:
121
- token = self.pipeline_kwargs.get("token", None)
122
- self.tokenizer = AutoTokenizer.from_pretrained(self.model_name_or_path, token=token)
123
- self.model = AutoModelForTokenClassification.from_pretrained(self.model_name_or_path, token=token)
124
-
125
- pipeline_params: dict[str, Any] = {
126
- "task": "ner",
127
- "model": self.model,
128
- "tokenizer": self.tokenizer,
129
- "aggregation_strategy": "simple",
130
- }
131
- pipeline_params.update({k: v for k, v in self.pipeline_kwargs.items() if k not in pipeline_params})
132
- self.device.update_hf_kwargs(pipeline_params, overwrite=False)
133
- self.pipeline = pipeline(**pipeline_params)
134
- self._warmed_up = True
116
+ pipeline_kwargs = _with_hf_token(self.pipeline_kwargs, self.token)
117
+ pipeline_kwargs.setdefault("aggregation_strategy", "simple")
118
+ self.pipeline = pipeline(**pipeline_kwargs)
135
119
  except Exception as e:
136
120
  msg = f"{self.__class__.__name__} failed to initialize."
137
121
  raise ComponentError(msg) from e
@@ -150,8 +134,8 @@ class TransformersNamedEntityExtractor:
150
134
  :raises ComponentError:
151
135
  If the model fails to process a document.
152
136
  """
153
- if not self._warmed_up:
154
- self.warm_up()
137
+ self.warm_up()
138
+ assert self.pipeline is not None # noqa: S101
155
139
 
156
140
  texts = [doc.content if doc.content is not None else "" for doc in documents]
157
141
  annotations = self._annotate(texts, batch_size=batch_size)
@@ -182,11 +166,10 @@ class TransformersNamedEntityExtractor:
182
166
  :returns:
183
167
  NER annotations.
184
168
  """
185
- if not self.initialized:
169
+ if self.pipeline is None:
186
170
  msg = "NER model was not initialized - Did you call `warm_up()`?"
187
171
  raise ComponentError(msg)
188
172
 
189
- assert self.pipeline is not None # noqa: S101
190
173
  outputs = self.pipeline(texts, batch_size=batch_size)
191
174
  return [
192
175
  [
@@ -246,7 +229,7 @@ class TransformersNamedEntityExtractor:
246
229
  """
247
230
  Returns if the extractor is ready to annotate text.
248
231
  """
249
- return (self.tokenizer is not None and self.model is not None) or self.pipeline is not None
232
+ return self.pipeline is not None
250
233
 
251
234
  @classmethod
252
235
  def get_stored_annotations(cls, document: Document) -> list[NamedEntityAnnotation] | None:
@@ -28,20 +28,19 @@ from haystack.tools.utils import warm_up_tools
28
28
  from haystack.utils import ComponentDevice, Secret, deserialize_callable, serialize_callable
29
29
  from haystack.utils.hf import convert_message_to_hf_format, deserialize_hf_model_kwargs, serialize_hf_model_kwargs
30
30
  from huggingface_hub import model_info
31
- from packaging.version import Version
32
31
 
33
- import transformers
34
32
  from haystack_integrations.common.transformers.utils import (
35
33
  _AsyncHFTokenStreamingHandler,
36
34
  _HFTokenStreamingHandler,
37
35
  _StopWordsCriteria,
36
+ _with_hf_token,
38
37
  )
39
38
  from transformers import Pipeline as HfPipeline
40
39
  from transformers import PreTrainedTokenizer, PreTrainedTokenizerFast, StoppingCriteriaList, pipeline
41
40
 
42
41
  logger = logging.getLogger(__name__)
43
42
 
44
- PIPELINE_SUPPORTED_TASKS = ["text-generation", "text2text-generation", "image-text-to-text"]
43
+ PIPELINE_SUPPORTED_TASKS = ["text-generation", "image-text-to-text"]
45
44
 
46
45
  DEFAULT_TOOL_PATTERN = (
47
46
  r"(?:<tool_call>)?"
@@ -119,7 +118,7 @@ class TransformersChatGenerator:
119
118
  def __init__(
120
119
  self,
121
120
  model: str = "Qwen/Qwen3-0.6B",
122
- task: Literal["text-generation", "text2text-generation", "image-text-to-text"] | None = None,
121
+ task: Literal["text-generation", "image-text-to-text"] | None = None,
123
122
  device: ComponentDevice | None = None,
124
123
  token: Secret | None = Secret.from_env_var(["HF_API_TOKEN", "HF_TOKEN"], strict=False),
125
124
  chat_template: str | None = None,
@@ -143,8 +142,6 @@ class TransformersChatGenerator:
143
142
  If the model is specified in `huggingface_pipeline_kwargs`, this parameter is ignored.
144
143
  :param task: The task for the Hugging Face pipeline. Possible options:
145
144
  - `text-generation`: Supported by decoder models, like GPT.
146
- - `text2text-generation`: Deprecated as of Transformers v5; use `text-generation` instead.
147
- Previously supported by encoder-decoder models such as T5.
148
145
  - `image-text-to-text`: Supported by vision-language models.
149
146
  If the task is specified in `huggingface_pipeline_kwargs`, this parameter is ignored.
150
147
  If not specified, the component calls the Hugging Face API to infer the task from the model name.
@@ -192,32 +189,17 @@ class TransformersChatGenerator:
192
189
  generation_kwargs = generation_kwargs or {}
193
190
 
194
191
  self.token = token
195
- token = token.resolve_value() if token else None
196
192
 
197
193
  # check if the huggingface_pipeline_kwargs contain the essential parameters
198
194
  # otherwise, populate them with values from other init parameters
199
195
  huggingface_pipeline_kwargs.setdefault("model", model)
200
- huggingface_pipeline_kwargs.setdefault("token", token)
201
196
 
202
197
  device = ComponentDevice.resolve_device(device)
203
198
  device.update_hf_kwargs(huggingface_pipeline_kwargs, overwrite=False)
204
199
 
205
- # task identification and validation
206
- if task is None:
207
- if "task" in huggingface_pipeline_kwargs:
208
- task = huggingface_pipeline_kwargs["task"]
209
- elif isinstance(huggingface_pipeline_kwargs["model"], str):
210
- task = model_info(
211
- huggingface_pipeline_kwargs["model"], token=huggingface_pipeline_kwargs["token"]
212
- ).pipeline_tag # type: ignore[assignment] # we'll check below if task is in supported tasks
213
-
214
- if task not in PIPELINE_SUPPORTED_TASKS:
215
- msg = f"Task '{task}' is not supported. The supported tasks are: {', '.join(PIPELINE_SUPPORTED_TASKS)}."
216
- raise ValueError(msg)
217
- if task == "text2text-generation" and Version(transformers.__version__) >= Version("5.0.0"):
218
- msg = "Task 'text2text-generation' is not supported with transformers v5 or higher."
219
- raise ValueError(msg)
220
- huggingface_pipeline_kwargs["task"] = task
200
+ task = task or huggingface_pipeline_kwargs.get("task")
201
+ if task is not None:
202
+ huggingface_pipeline_kwargs["task"] = task
221
203
 
222
204
  # if not specified, set return_full_text to False for text-generation
223
205
  # only generated text is returned (excluding prompt)
@@ -244,26 +226,7 @@ class TransformersChatGenerator:
244
226
  self.enable_thinking = enable_thinking
245
227
 
246
228
  self._owns_executor = async_executor is None
247
- self.executor = (
248
- ThreadPoolExecutor(thread_name_prefix=f"async-TransformersChatGenerator-executor-{id(self)}", max_workers=1)
249
- if async_executor is None
250
- else async_executor
251
- )
252
- self._is_warmed_up = False
253
-
254
- def __del__(self) -> None:
255
- """
256
- Cleanup when the instance is being destroyed.
257
- """
258
- if hasattr(self, "_owns_executor") and self._owns_executor and hasattr(self, "executor"):
259
- self.executor.shutdown(wait=True)
260
-
261
- def shutdown(self) -> None:
262
- """
263
- Explicitly shutdown the executor if we own it.
264
- """
265
- if self._owns_executor:
266
- self.executor.shutdown(wait=True)
229
+ self.executor = async_executor
267
230
 
268
231
  def _get_telemetry_data(self) -> dict[str, Any]:
269
232
  """
@@ -277,18 +240,34 @@ class TransformersChatGenerator:
277
240
  """
278
241
  Initializes the component and warms up tools if provided.
279
242
  """
280
- if self._is_warmed_up:
281
- return
282
-
283
- # Initialize the pipeline
284
243
  if self.pipeline is None:
285
- self.pipeline = pipeline(**self.huggingface_pipeline_kwargs)
286
-
287
- # Warm up tools
288
- if self.tools:
289
- warm_up_tools(self.tools)
244
+ pipeline_kwargs = _with_hf_token(self.huggingface_pipeline_kwargs, self.token)
245
+ task = pipeline_kwargs.get("task")
246
+ if task is None and isinstance(pipeline_kwargs["model"], str):
247
+ task = model_info(pipeline_kwargs["model"], token=pipeline_kwargs["token"]).pipeline_tag
248
+
249
+ if task not in PIPELINE_SUPPORTED_TASKS:
250
+ msg = f"Task '{task}' is not supported. The supported tasks are: {', '.join(PIPELINE_SUPPORTED_TASKS)}."
251
+ raise ValueError(msg)
252
+ pipeline_kwargs["task"] = task
253
+
254
+ hf_pipeline = pipeline(**pipeline_kwargs)
255
+ if self.tools:
256
+ warm_up_tools(self.tools)
257
+ self.pipeline = hf_pipeline
258
+
259
+ if self._owns_executor and self.executor is None:
260
+ self.executor = ThreadPoolExecutor(
261
+ thread_name_prefix=f"async-TransformersChatGenerator-executor-{id(self)}", max_workers=1
262
+ )
290
263
 
291
- self._is_warmed_up = True
264
+ def close(self) -> None:
265
+ """
266
+ Close the executor owned by the component.
267
+ """
268
+ if self._owns_executor and self.executor is not None:
269
+ self.executor.shutdown(wait=True)
270
+ self.executor = None
292
271
 
293
272
  def to_dict(self) -> dict[str, Any]:
294
273
  """
@@ -353,15 +332,17 @@ class TransformersChatGenerator:
353
332
 
354
333
  :param messages: A list of ChatMessage objects representing the input messages. If a string is provided,
355
334
  it is converted to a list containing a ChatMessage with user role.
356
- :param generation_kwargs: Additional keyword arguments for text generation.
335
+ :param generation_kwargs: Additional keyword arguments for text generation. These are merged per key with
336
+ the `generation_kwargs` passed at initialization: keys provided here take precedence, keys set only at
337
+ initialization are kept.
357
338
  :param streaming_callback: An optional callable for handling streaming responses.
358
339
  :param tools: A list of Tool and/or Toolset objects, or a single Toolset for which the model can prepare calls.
359
340
  If set, it will override the `tools` parameter provided during initialization.
360
341
  :returns: A dictionary with the following keys:
361
342
  - `replies`: A list containing the generated responses as ChatMessage instances.
362
343
  """
363
- if self.pipeline is None:
364
- self.warm_up()
344
+ self.warm_up()
345
+ assert self.pipeline is not None # noqa: S101
365
346
 
366
347
  messages = _normalize_messages(messages)
367
348
 
@@ -381,8 +362,6 @@ class TransformersChatGenerator:
381
362
  component_info=ComponentInfo.from_component(self),
382
363
  )
383
364
 
384
- # We know it's not None because we check it in _prepare_inputs
385
- assert self.pipeline is not None # noqa: S101
386
365
  # Generate responses
387
366
  output = self.pipeline(prepared_inputs["prepared_prompt"], **prepared_inputs["generation_kwargs"])
388
367
 
@@ -472,15 +451,18 @@ class TransformersChatGenerator:
472
451
  and return values but can be used with `await` in an async code.
473
452
 
474
453
  :param messages: A list of ChatMessage objects representing the input messages.
475
- :param generation_kwargs: Additional keyword arguments for text generation.
454
+ :param generation_kwargs: Additional keyword arguments for text generation. These are merged per key with
455
+ the `generation_kwargs` passed at initialization: keys provided here take precedence, keys set only at
456
+ initialization are kept.
476
457
  :param streaming_callback: An optional callable for handling streaming responses.
477
458
  :param tools: A list of Tool and/or Toolset objects, or a single Toolset for which the model can prepare calls.
478
459
  If set, it will override the `tools` parameter provided during initialization.
479
460
  :returns: A dictionary with the following keys:
480
461
  - `replies`: A list containing the generated responses as ChatMessage instances.
481
462
  """
482
- if self.pipeline is None:
483
- self.warm_up()
463
+ self.warm_up()
464
+ assert self.pipeline is not None # noqa: S101
465
+ assert self.executor is not None # noqa: S101
484
466
 
485
467
  messages = _normalize_messages(messages)
486
468
 
@@ -555,6 +537,8 @@ class TransformersChatGenerator:
555
537
  :returns: A dictionary containing the prepared prompt, tokenizer, generation kwargs, and tools.
556
538
  :raises ValueError: If both tools and streaming_callback are provided.
557
539
  """
540
+ assert self.pipeline is not None # noqa: S101
541
+
558
542
  tools = tools or self.tools
559
543
  if tools and streaming_callback is not None:
560
544
  msg = "Using tools and streaming at the same time is not supported. Please choose one."
@@ -562,11 +546,12 @@ class TransformersChatGenerator:
562
546
  flat_tools = flatten_tools_or_toolsets(tools)
563
547
  _check_duplicate_tool_names(flat_tools)
564
548
 
565
- # mypy doesn't know this is set in warm_up
566
- tokenizer = self.pipeline.tokenizer # type: ignore[union-attr]
549
+ tokenizer = self.pipeline.tokenizer
567
550
 
568
551
  # Check and update generation parameters
569
552
  generation_kwargs = {**self.generation_kwargs, **(generation_kwargs or {})}
553
+ if self.pipeline.task == "text-generation":
554
+ generation_kwargs.setdefault("return_full_text", False)
570
555
 
571
556
  # If streaming_callback is provided, ensure that num_return_sequences is set to 1
572
557
  if streaming_callback:
@@ -588,7 +573,7 @@ class TransformersChatGenerator:
588
573
  _StopWordsCriteria(
589
574
  tokenizer, # type: ignore[arg-type]
590
575
  stop_words,
591
- self.pipeline.device, # type: ignore[union-attr]
576
+ self.pipeline.device,
592
577
  )
593
578
  if stop_words
594
579
  else None
@@ -176,21 +176,23 @@ class TransformersExtractiveReader:
176
176
  """
177
177
  Initializes the component.
178
178
  """
179
+ if self.model is not None:
180
+ return
181
+
179
182
  # Take the first device used by `accelerate`. Needed to pass inputs from the tokenizer to the correct device.
180
- if self.model is None:
181
- self.model = AutoModelForQuestionAnswering.from_pretrained(
182
- self.model_name_or_path, token=self.token.resolve_value() if self.token else None, **self.model_kwargs
183
- )
184
- self.tokenizer = AutoTokenizer.from_pretrained(
185
- self.model_name_or_path, token=self.token.resolve_value() if self.token else None
186
- )
187
- assert self.model is not None # noqa: S101 # mypy doesn't know this is set in the line above
188
- # hf_device_map appears to only be set now when mixed devices are actually used.
189
- # So if it's missing then we can use the device attribute which is set even for single-device models.
190
- if hf_device_map := getattr(self.model, "hf_device_map", None):
191
- self.device = ComponentDevice.from_multiple(device_map=DeviceMap.from_hf(hf_device_map))
192
- else:
193
- self.device = ComponentDevice.from_single(Device.from_str(str(self.model.device)))
183
+ token = self.token.resolve_value() if self.token else None
184
+ model = AutoModelForQuestionAnswering.from_pretrained(self.model_name_or_path, token=token, **self.model_kwargs)
185
+ tokenizer = AutoTokenizer.from_pretrained(self.model_name_or_path, token=token)
186
+ # hf_device_map appears to only be set now when mixed devices are actually used.
187
+ # So if it's missing then we can use the device attribute which is set even for single-device models.
188
+ if hf_device_map := getattr(model, "hf_device_map", None):
189
+ device = ComponentDevice.from_multiple(device_map=DeviceMap.from_hf(hf_device_map))
190
+ else:
191
+ device = ComponentDevice.from_single(Device.from_str(str(model.device)))
192
+
193
+ self.model = model
194
+ self.tokenizer = tokenizer
195
+ self.device = device
194
196
 
195
197
  @staticmethod
196
198
  def _flatten_documents(
@@ -314,8 +316,10 @@ class TransformersExtractiveReader:
314
316
  # But we shouldn't have special tokens in the answers at this point
315
317
  # The whole span is given by the start of the start_token (index 0)
316
318
  # and the end of the end token (index 1)
317
- s_char_spans.append(encoding.token_to_chars(start_token)[0])
318
- e_char_spans.append(encoding.token_to_chars(end_token)[1])
319
+ # `type: ignore[index]` because tokenizers>=0.23.1 types the return as
320
+ # `tuple[int, int] | None`; the `None` case cannot occur here per the above
321
+ s_char_spans.append(encoding.token_to_chars(start_token)[0]) # type: ignore[index]
322
+ e_char_spans.append(encoding.token_to_chars(end_token)[1]) # type: ignore[index]
319
323
  start_candidates_tokens_to_chars.append(s_char_spans)
320
324
  end_candidates_tokens_to_chars.append(e_char_spans)
321
325
  valid_candidates_values.append(candidates_values[i][valid])
@@ -581,8 +585,10 @@ class TransformersExtractiveReader:
581
585
  :returns:
582
586
  List of answers sorted by (desc.) answer score.
583
587
  """
584
- if self.model is None:
585
- self.warm_up()
588
+ self.warm_up()
589
+ assert self.model is not None # noqa: S101
590
+ assert self.tokenizer is not None # noqa: S101
591
+ assert self.device is not None # noqa: S101
586
592
 
587
593
  if not documents:
588
594
  return {"answers": []}