transformers-haystack 0.1.0__tar.gz → 0.3.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.
- transformers_haystack-0.3.0/CHANGELOG.md +28 -0
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/PKG-INFO +2 -2
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/pyproject.toml +2 -1
- {transformers_haystack-0.1.0/src/haystack_integrations/components → transformers_haystack-0.3.0/src/haystack_integrations}/common/transformers/utils.py +8 -4
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/src/haystack_integrations/components/classifiers/transformers/zero_shot_document_classifier.py +1 -1
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/src/haystack_integrations/components/extractors/transformers/named_entity_extractor.py +1 -1
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/src/haystack_integrations/components/generators/transformers/chat/chat_generator.py +7 -3
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/src/haystack_integrations/components/readers/transformers/extractive_reader.py +1 -1
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/src/haystack_integrations/components/routers/transformers/text_router.py +1 -1
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/src/haystack_integrations/components/routers/transformers/zero_shot_text_router.py +9 -3
- transformers_haystack-0.3.0/tests/conftest.py +39 -0
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/tests/test_chat_generator.py +60 -49
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/tests/test_extractive_reader.py +18 -37
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/tests/test_named_entity_extractor.py +9 -9
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/tests/test_text_router.py +16 -10
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/tests/test_utils.py +1 -1
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/tests/test_zero_shot_document_classifier.py +8 -6
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/tests/test_zero_shot_text_router.py +15 -6
- transformers_haystack-0.1.0/tests/conftest.py +0 -26
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/.gitignore +0 -0
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/LICENSE.txt +0 -0
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/README.md +0 -0
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/pydoc/config_docusaurus.yml +0 -0
- {transformers_haystack-0.1.0/src/haystack_integrations/components/classifiers → transformers_haystack-0.3.0/src/haystack_integrations/common}/py.typed +0 -0
- {transformers_haystack-0.1.0/src/haystack_integrations/components → transformers_haystack-0.3.0/src/haystack_integrations}/common/transformers/__init__.py +0 -0
- {transformers_haystack-0.1.0/src/haystack_integrations/components/common → transformers_haystack-0.3.0/src/haystack_integrations/components/classifiers}/py.typed +0 -0
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/src/haystack_integrations/components/classifiers/transformers/__init__.py +0 -0
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/src/haystack_integrations/components/extractors/py.typed +0 -0
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/src/haystack_integrations/components/extractors/transformers/__init__.py +0 -0
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/src/haystack_integrations/components/generators/py.typed +0 -0
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/src/haystack_integrations/components/generators/transformers/__init__.py +0 -0
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/src/haystack_integrations/components/generators/transformers/chat/__init__.py +0 -0
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/src/haystack_integrations/components/readers/py.typed +0 -0
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/src/haystack_integrations/components/readers/transformers/__init__.py +0 -0
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/src/haystack_integrations/components/routers/py.typed +0 -0
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/src/haystack_integrations/components/routers/transformers/__init__.py +0 -0
- {transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/tests/__init__.py +0 -0
|
@@ -0,0 +1,28 @@
|
|
|
1
|
+
# Changelog
|
|
2
|
+
|
|
3
|
+
## [integrations/transformers-v0.2.0] - 2026-07-06
|
|
4
|
+
|
|
5
|
+
### 📚 Documentation
|
|
6
|
+
|
|
7
|
+
- Replace old haystack core imports with haystack_integrations paths (#3545)
|
|
8
|
+
|
|
9
|
+
### 🧪 Testing
|
|
10
|
+
|
|
11
|
+
- Improve del_hf_env_vars fixture (#3428)
|
|
12
|
+
- Trust test modules under Haystack 3.0's deserialization allowlist (#3537)
|
|
13
|
+
- Make Tool/Agent serialization assertions version-agnostic for Haystack 2.x/3.x (#3533)
|
|
14
|
+
- Force Transformers and Sentence Transformers integration tests to run on CPU (#3550)
|
|
15
|
+
|
|
16
|
+
### 🧹 Chores
|
|
17
|
+
|
|
18
|
+
- Improve consistency of integrations folder structure (#3430)
|
|
19
|
+
- Support sync streaming callbacks in async contexts for Haystack 2.x/3.x compatibility (#3534)
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
## [integrations/transformers-v0.1.0] - 2026-06-08
|
|
23
|
+
|
|
24
|
+
### 🚀 Features
|
|
25
|
+
|
|
26
|
+
- Move Transformers components from Haystack (#3409)
|
|
27
|
+
|
|
28
|
+
<!-- generated by git-cliff -->
|
|
@@ -1,6 +1,6 @@
|
|
|
1
|
-
Metadata-Version: 2.
|
|
1
|
+
Metadata-Version: 2.5
|
|
2
2
|
Name: transformers-haystack
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.3.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
|
|
@@ -66,7 +66,7 @@ integration = 'pytest -m "integration" {args:tests}'
|
|
|
66
66
|
all = 'pytest {args:tests}'
|
|
67
67
|
unit-cov-retry = 'pytest --cov=haystack_integrations --reruns 3 --reruns-delay 30 -x -m "not integration" {args:tests}'
|
|
68
68
|
integration-cov-append-retry = 'pytest --cov=haystack_integrations --cov-append --reruns 3 --reruns-delay 30 -x -m "integration" {args:tests}'
|
|
69
|
-
types = """mypy -p haystack_integrations.
|
|
69
|
+
types = """mypy -p haystack_integrations.common.transformers \
|
|
70
70
|
-p haystack_integrations.components.classifiers.transformers \
|
|
71
71
|
-p haystack_integrations.components.extractors.transformers \
|
|
72
72
|
-p haystack_integrations.components.generators.transformers \
|
|
@@ -140,6 +140,7 @@ ignore = [
|
|
|
140
140
|
"PLR0912",
|
|
141
141
|
"PLR0913",
|
|
142
142
|
"PLR0915",
|
|
143
|
+
"PLR0917",
|
|
143
144
|
# Allow `Any` type - used legitimately for dynamic types and SDK boundaries
|
|
144
145
|
"ANN401",
|
|
145
146
|
]
|
|
@@ -4,11 +4,12 @@
|
|
|
4
4
|
|
|
5
5
|
import asyncio
|
|
6
6
|
import copy
|
|
7
|
+
import inspect
|
|
7
8
|
from typing import Any
|
|
8
9
|
|
|
9
10
|
import torch
|
|
10
11
|
from haystack import logging
|
|
11
|
-
from haystack.dataclasses import
|
|
12
|
+
from haystack.dataclasses import ComponentInfo, StreamingCallbackT, StreamingChunk, SyncStreamingCallbackT
|
|
12
13
|
from haystack.utils.auth import Secret
|
|
13
14
|
from haystack.utils.device import ComponentDevice
|
|
14
15
|
from huggingface_hub import model_info
|
|
@@ -198,7 +199,7 @@ class _AsyncHFTokenStreamingHandler(TextStreamer):
|
|
|
198
199
|
Async streaming handler for TransformersChatGenerator.
|
|
199
200
|
|
|
200
201
|
Note: This is a helper class for TransformersChatGenerator enabling
|
|
201
|
-
async streaming of generated text via Haystack
|
|
202
|
+
async streaming of generated text via Haystack StreamingCallbackT callbacks.
|
|
202
203
|
|
|
203
204
|
Do not use this class directly.
|
|
204
205
|
"""
|
|
@@ -206,7 +207,7 @@ class _AsyncHFTokenStreamingHandler(TextStreamer):
|
|
|
206
207
|
def __init__(
|
|
207
208
|
self,
|
|
208
209
|
tokenizer: PreTrainedTokenizerBase,
|
|
209
|
-
stream_handler:
|
|
210
|
+
stream_handler: StreamingCallbackT,
|
|
210
211
|
stop_words: list[str] | None = None,
|
|
211
212
|
component_info: ComponentInfo | None = None,
|
|
212
213
|
) -> None:
|
|
@@ -228,7 +229,10 @@ class _AsyncHFTokenStreamingHandler(TextStreamer):
|
|
|
228
229
|
while True:
|
|
229
230
|
try:
|
|
230
231
|
chunk = await self._queue.get()
|
|
231
|
-
await
|
|
232
|
+
# sync callbacks are allowed in async contexts with Haystack >= 3.0, so only await async ones
|
|
233
|
+
callback_result = self.token_handler(chunk)
|
|
234
|
+
if inspect.isawaitable(callback_result):
|
|
235
|
+
await callback_result
|
|
232
236
|
self._queue.task_done()
|
|
233
237
|
except asyncio.CancelledError:
|
|
234
238
|
break
|
|
@@ -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.
|
|
12
|
+
from haystack_integrations.common.transformers.utils import _resolve_hf_pipeline_kwargs
|
|
13
13
|
from transformers import Pipeline as HfPipeline
|
|
14
14
|
from transformers import pipeline
|
|
15
15
|
|
|
@@ -10,7 +10,7 @@ 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.
|
|
13
|
+
from haystack_integrations.common.transformers.utils import _resolve_hf_pipeline_kwargs
|
|
14
14
|
from transformers import AutoModelForTokenClassification, AutoTokenizer, pipeline
|
|
15
15
|
from transformers import Pipeline as HfPipeline
|
|
16
16
|
|
|
@@ -31,7 +31,7 @@ from huggingface_hub import model_info
|
|
|
31
31
|
from packaging.version import Version
|
|
32
32
|
|
|
33
33
|
import transformers
|
|
34
|
-
from haystack_integrations.
|
|
34
|
+
from haystack_integrations.common.transformers.utils import (
|
|
35
35
|
_AsyncHFTokenStreamingHandler,
|
|
36
36
|
_HFTokenStreamingHandler,
|
|
37
37
|
_StopWordsCriteria,
|
|
@@ -353,7 +353,9 @@ class TransformersChatGenerator:
|
|
|
353
353
|
|
|
354
354
|
:param messages: A list of ChatMessage objects representing the input messages. If a string is provided,
|
|
355
355
|
it is converted to a list containing a ChatMessage with user role.
|
|
356
|
-
:param generation_kwargs: Additional keyword arguments for text generation.
|
|
356
|
+
:param generation_kwargs: Additional keyword arguments for text generation. These are merged per key with
|
|
357
|
+
the `generation_kwargs` passed at initialization: keys provided here take precedence, keys set only at
|
|
358
|
+
initialization are kept.
|
|
357
359
|
:param streaming_callback: An optional callable for handling streaming responses.
|
|
358
360
|
:param tools: A list of Tool and/or Toolset objects, or a single Toolset for which the model can prepare calls.
|
|
359
361
|
If set, it will override the `tools` parameter provided during initialization.
|
|
@@ -472,7 +474,9 @@ class TransformersChatGenerator:
|
|
|
472
474
|
and return values but can be used with `await` in an async code.
|
|
473
475
|
|
|
474
476
|
:param messages: A list of ChatMessage objects representing the input messages.
|
|
475
|
-
:param generation_kwargs: Additional keyword arguments for text generation.
|
|
477
|
+
:param generation_kwargs: Additional keyword arguments for text generation. These are merged per key with
|
|
478
|
+
the `generation_kwargs` passed at initialization: keys provided here take precedence, keys set only at
|
|
479
|
+
initialization are kept.
|
|
476
480
|
:param streaming_callback: An optional callable for handling streaming responses.
|
|
477
481
|
:param tools: A list of Tool and/or Toolset objects, or a single Toolset for which the model can prepare calls.
|
|
478
482
|
If set, it will override the `tools` parameter provided during initialization.
|
|
@@ -14,7 +14,7 @@ from haystack.utils import ComponentDevice, Device, DeviceMap, Secret
|
|
|
14
14
|
from haystack.utils.hf import deserialize_hf_model_kwargs, serialize_hf_model_kwargs
|
|
15
15
|
from tokenizers import Encoding
|
|
16
16
|
|
|
17
|
-
from haystack_integrations.
|
|
17
|
+
from haystack_integrations.common.transformers.utils import _resolve_hf_device_map
|
|
18
18
|
from transformers import AutoModelForQuestionAnswering, AutoTokenizer
|
|
19
19
|
|
|
20
20
|
logger = logging.getLogger(__name__)
|
|
@@ -8,7 +8,7 @@ from haystack import component, default_from_dict, default_to_dict
|
|
|
8
8
|
from haystack.utils import ComponentDevice, Secret
|
|
9
9
|
from haystack.utils.hf import deserialize_hf_model_kwargs, serialize_hf_model_kwargs
|
|
10
10
|
|
|
11
|
-
from haystack_integrations.
|
|
11
|
+
from haystack_integrations.common.transformers.utils import _resolve_hf_pipeline_kwargs
|
|
12
12
|
from transformers import AutoConfig, Pipeline, pipeline
|
|
13
13
|
|
|
14
14
|
|
|
@@ -8,7 +8,7 @@ from haystack import component, default_from_dict, default_to_dict
|
|
|
8
8
|
from haystack.utils import ComponentDevice, Secret
|
|
9
9
|
from haystack.utils.hf import deserialize_hf_model_kwargs, serialize_hf_model_kwargs
|
|
10
10
|
|
|
11
|
-
from haystack_integrations.
|
|
11
|
+
from haystack_integrations.common.transformers.utils import _resolve_hf_pipeline_kwargs
|
|
12
12
|
from transformers import Pipeline as HfPipeline
|
|
13
13
|
from transformers import pipeline
|
|
14
14
|
|
|
@@ -24,7 +24,9 @@ class TransformersZeroShotTextRouter:
|
|
|
24
24
|
|
|
25
25
|
```python
|
|
26
26
|
from haystack import Document
|
|
27
|
-
|
|
27
|
+
# Requires: pip install sentence-transformers-haystack
|
|
28
|
+
from haystack_integrations.components.embedders.sentence_transformers import SentenceTransformersTextEmbedder
|
|
29
|
+
from haystack_integrations.components.embedders.sentence_transformers import SentenceTransformersDocumentEmbedder
|
|
28
30
|
from haystack.components.retrievers import InMemoryEmbeddingRetriever
|
|
29
31
|
from haystack.core.pipeline import Pipeline
|
|
30
32
|
from haystack.document_stores.in_memory import InMemoryDocumentStore
|
|
@@ -154,7 +156,11 @@ class TransformersZeroShotTextRouter:
|
|
|
154
156
|
Dictionary with serialized data.
|
|
155
157
|
"""
|
|
156
158
|
serialization_dict = default_to_dict(
|
|
157
|
-
self,
|
|
159
|
+
self,
|
|
160
|
+
labels=self.labels,
|
|
161
|
+
multi_label=self.multi_label,
|
|
162
|
+
huggingface_pipeline_kwargs=self.huggingface_pipeline_kwargs,
|
|
163
|
+
token=self.token,
|
|
158
164
|
)
|
|
159
165
|
|
|
160
166
|
huggingface_pipeline_kwargs = serialization_dict["init_parameters"]["huggingface_pipeline_kwargs"]
|
|
@@ -0,0 +1,39 @@
|
|
|
1
|
+
# SPDX-FileCopyrightText: 2022-present deepset GmbH <info@deepset.ai>
|
|
2
|
+
#
|
|
3
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
4
|
+
|
|
5
|
+
import os
|
|
6
|
+
|
|
7
|
+
import pytest
|
|
8
|
+
from haystack.document_stores.in_memory import InMemoryDocumentStore
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
@pytest.fixture()
|
|
12
|
+
def del_hf_env_vars_if_empty(monkeypatch):
|
|
13
|
+
"""
|
|
14
|
+
Delete Hugging Face environment variables for tests if empty.
|
|
15
|
+
|
|
16
|
+
Prevents passing empty tokens to Hugging Face, which would cause API calls to fail.
|
|
17
|
+
This is particularly relevant for PRs opened from forks, where secrets are not available
|
|
18
|
+
and empty environment variables might be set instead of being removed.
|
|
19
|
+
|
|
20
|
+
See https://github.com/deepset-ai/haystack/issues/8811 for more details.
|
|
21
|
+
"""
|
|
22
|
+
for var in ("HF_API_TOKEN", "HF_TOKEN"):
|
|
23
|
+
if not os.environ.get(var, "").strip():
|
|
24
|
+
monkeypatch.delenv(var, raising=False)
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
@pytest.fixture()
|
|
28
|
+
def in_memory_doc_store():
|
|
29
|
+
return InMemoryDocumentStore()
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
@pytest.fixture(autouse=True)
|
|
33
|
+
def allow_deserialization_of_test_modules(monkeypatch):
|
|
34
|
+
"""
|
|
35
|
+
haystack-ai >= 3.0 refuses to deserialize classes and callables from modules outside its
|
|
36
|
+
trusted-module allowlist. Tools and callbacks defined in the test modules live outside that
|
|
37
|
+
allowlist, so trust them explicitly; haystack-ai 2.x ignores this environment variable.
|
|
38
|
+
"""
|
|
39
|
+
monkeypatch.setenv("HAYSTACK_DESERIALIZATION_ALLOWLIST", "tests,test_*")
|
|
@@ -206,20 +206,10 @@ class TestTransformersChatGenerator:
|
|
|
206
206
|
assert init_params["streaming_callback"] is None
|
|
207
207
|
assert init_params["chat_template"] == "irrelevant"
|
|
208
208
|
assert init_params["enable_thinking"] is True
|
|
209
|
-
|
|
210
|
-
|
|
211
|
-
|
|
212
|
-
|
|
213
|
-
"inputs_from_state": None,
|
|
214
|
-
"name": "weather",
|
|
215
|
-
"outputs_to_state": None,
|
|
216
|
-
"outputs_to_string": None,
|
|
217
|
-
"description": "useful to determine the weather in a given location",
|
|
218
|
-
"parameters": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]},
|
|
219
|
-
"function": "tests.test_chat_generator.get_weather",
|
|
220
|
-
},
|
|
221
|
-
}
|
|
222
|
-
]
|
|
209
|
+
|
|
210
|
+
# deserializing the serialized component must reproduce the original tools
|
|
211
|
+
loaded = TransformersChatGenerator.from_dict(result)
|
|
212
|
+
assert loaded.tools == tools
|
|
223
213
|
|
|
224
214
|
def test_from_dict(self, model_info_mock, tools):
|
|
225
215
|
generator = TransformersChatGenerator(
|
|
@@ -256,7 +246,7 @@ class TestTransformersChatGenerator:
|
|
|
256
246
|
}
|
|
257
247
|
|
|
258
248
|
@patch("haystack_integrations.components.generators.transformers.chat.chat_generator.pipeline")
|
|
259
|
-
def test_warm_up(self, pipeline_mock,
|
|
249
|
+
def test_warm_up(self, pipeline_mock, del_hf_env_vars_if_empty):
|
|
260
250
|
generator = TransformersChatGenerator(
|
|
261
251
|
model="mistralai/Mistral-7B-Instruct-v0.2", task="text-generation", device=ComponentDevice.from_str("cpu")
|
|
262
252
|
)
|
|
@@ -266,11 +256,14 @@ class TestTransformersChatGenerator:
|
|
|
266
256
|
generator.warm_up()
|
|
267
257
|
|
|
268
258
|
pipeline_mock.assert_called_once_with(
|
|
269
|
-
model="mistralai/Mistral-7B-Instruct-v0.2",
|
|
259
|
+
model="mistralai/Mistral-7B-Instruct-v0.2",
|
|
260
|
+
task="text-generation",
|
|
261
|
+
token=generator.token.resolve_value(),
|
|
262
|
+
device="cpu",
|
|
270
263
|
)
|
|
271
264
|
|
|
272
265
|
@patch("haystack_integrations.components.generators.transformers.chat.chat_generator.pipeline")
|
|
273
|
-
def test_warm_up_with_tools(self, pipeline_mock,
|
|
266
|
+
def test_warm_up_with_tools(self, pipeline_mock, del_hf_env_vars_if_empty):
|
|
274
267
|
"""Test that warm_up() calls warm_up on tools and is idempotent."""
|
|
275
268
|
|
|
276
269
|
# Create a mock tool that tracks if warm_up() was called
|
|
@@ -324,7 +317,7 @@ class TestTransformersChatGenerator:
|
|
|
324
317
|
pipeline_mock.assert_called_once()
|
|
325
318
|
|
|
326
319
|
@patch("haystack_integrations.components.generators.transformers.chat.chat_generator.pipeline")
|
|
327
|
-
def test_warm_up_with_no_tools(self, pipeline_mock,
|
|
320
|
+
def test_warm_up_with_no_tools(self, pipeline_mock, del_hf_env_vars_if_empty):
|
|
328
321
|
"""Test that warm_up() works when no tools are provided."""
|
|
329
322
|
|
|
330
323
|
generator = TransformersChatGenerator(
|
|
@@ -349,7 +342,7 @@ class TestTransformersChatGenerator:
|
|
|
349
342
|
pipeline_mock.assert_called_once()
|
|
350
343
|
|
|
351
344
|
@patch("haystack_integrations.components.generators.transformers.chat.chat_generator.pipeline")
|
|
352
|
-
def test_warm_up_with_multiple_tools(self, pipeline_mock,
|
|
345
|
+
def test_warm_up_with_multiple_tools(self, pipeline_mock, del_hf_env_vars_if_empty):
|
|
353
346
|
"""Test that warm_up() works with multiple tools."""
|
|
354
347
|
|
|
355
348
|
# Track warm_up calls
|
|
@@ -447,6 +440,19 @@ class TestTransformersChatGenerator:
|
|
|
447
440
|
assert chat_message.is_from(ChatRole.ASSISTANT)
|
|
448
441
|
assert chat_message.text == "Berlin is cool"
|
|
449
442
|
|
|
443
|
+
def test_run_with_generation_kwargs(self, model_info_mock, mock_pipeline_with_tokenizer, chat_messages):
|
|
444
|
+
generator = TransformersChatGenerator(
|
|
445
|
+
model="meta-llama/Llama-2-13b-chat-hf",
|
|
446
|
+
generation_kwargs={"max_new_tokens": 100, "temperature": 0.5},
|
|
447
|
+
)
|
|
448
|
+
generator.pipeline = mock_pipeline_with_tokenizer
|
|
449
|
+
|
|
450
|
+
generator.run(messages=chat_messages, generation_kwargs={"temperature": 0.9})
|
|
451
|
+
|
|
452
|
+
_, kwargs = generator.pipeline.call_args
|
|
453
|
+
assert kwargs["max_new_tokens"] == 100
|
|
454
|
+
assert kwargs["temperature"] == 0.9
|
|
455
|
+
|
|
450
456
|
def test_run_with_streaming_callback(self, model_info_mock, mock_pipeline_with_tokenizer, chat_messages):
|
|
451
457
|
# Define the streaming callback function
|
|
452
458
|
def streaming_callback_fn(chunk: StreamingChunk): ...
|
|
@@ -507,11 +513,15 @@ class TestTransformersChatGenerator:
|
|
|
507
513
|
|
|
508
514
|
@pytest.mark.integration
|
|
509
515
|
@pytest.mark.flaky(reruns=3, reruns_delay=10)
|
|
510
|
-
def test_live_run(self,
|
|
516
|
+
def test_live_run(self, del_hf_env_vars_if_empty):
|
|
511
517
|
"""Test live run with default behavior (no thinking)."""
|
|
512
518
|
messages = [ChatMessage.from_user("Please create a summary about the following topic: Climate change")]
|
|
513
519
|
|
|
514
|
-
llm = TransformersChatGenerator(
|
|
520
|
+
llm = TransformersChatGenerator(
|
|
521
|
+
model="Qwen/Qwen3-0.6B",
|
|
522
|
+
generation_kwargs={"max_new_tokens": 50},
|
|
523
|
+
device=ComponentDevice.from_str("cpu"),
|
|
524
|
+
)
|
|
515
525
|
|
|
516
526
|
result = llm.run(messages)
|
|
517
527
|
|
|
@@ -521,12 +531,15 @@ class TestTransformersChatGenerator:
|
|
|
521
531
|
|
|
522
532
|
@pytest.mark.integration
|
|
523
533
|
@pytest.mark.flaky(reruns=3, reruns_delay=10)
|
|
524
|
-
def test_live_run_thinking(self,
|
|
534
|
+
def test_live_run_thinking(self, del_hf_env_vars_if_empty):
|
|
525
535
|
"""Test live run with enable_thinking=True."""
|
|
526
536
|
messages = [ChatMessage.from_user("What is 2+2?")]
|
|
527
537
|
|
|
528
538
|
llm = TransformersChatGenerator(
|
|
529
|
-
model="Qwen/Qwen3-0.6B",
|
|
539
|
+
model="Qwen/Qwen3-0.6B",
|
|
540
|
+
generation_kwargs={"max_new_tokens": 450},
|
|
541
|
+
enable_thinking=True,
|
|
542
|
+
device=ComponentDevice.from_str("cpu"),
|
|
530
543
|
)
|
|
531
544
|
|
|
532
545
|
result = llm.run(messages)
|
|
@@ -705,6 +718,21 @@ class TestTransformersChatGeneratorAsync:
|
|
|
705
718
|
assert chat_message.text == "Berlin is cool"
|
|
706
719
|
generator.shutdown()
|
|
707
720
|
|
|
721
|
+
@pytest.mark.asyncio
|
|
722
|
+
async def test_run_async_with_generation_kwargs(self, model_info_mock, mock_pipeline_with_tokenizer, chat_messages):
|
|
723
|
+
generator = TransformersChatGenerator(
|
|
724
|
+
model="meta-llama/Llama-2-13b-chat-hf",
|
|
725
|
+
generation_kwargs={"max_new_tokens": 100, "temperature": 0.5},
|
|
726
|
+
)
|
|
727
|
+
generator.pipeline = mock_pipeline_with_tokenizer
|
|
728
|
+
|
|
729
|
+
await generator.run_async(messages=chat_messages, generation_kwargs={"temperature": 0.9})
|
|
730
|
+
|
|
731
|
+
_, kwargs = generator.pipeline.call_args
|
|
732
|
+
assert kwargs["max_new_tokens"] == 100
|
|
733
|
+
assert kwargs["temperature"] == 0.9
|
|
734
|
+
generator.shutdown()
|
|
735
|
+
|
|
708
736
|
@pytest.mark.asyncio
|
|
709
737
|
async def test_run_async_with_string_input(self, model_info_mock, mock_pipeline_with_tokenizer):
|
|
710
738
|
generator = TransformersChatGenerator(model="meta-llama/Llama-2-13b-chat-hf")
|
|
@@ -799,30 +827,10 @@ class TestTransformersChatGeneratorAsync:
|
|
|
799
827
|
generator.pipeline = mock_pipeline_with_tokenizer
|
|
800
828
|
data = generator.to_dict()
|
|
801
829
|
|
|
802
|
-
|
|
803
|
-
|
|
804
|
-
|
|
805
|
-
|
|
806
|
-
{
|
|
807
|
-
"type": "haystack.tools.tool.Tool",
|
|
808
|
-
"data": {
|
|
809
|
-
"name": "weather",
|
|
810
|
-
"description": "useful to determine the weather in a given location",
|
|
811
|
-
"parameters": {
|
|
812
|
-
"type": "object",
|
|
813
|
-
"properties": {"city": {"type": "string"}},
|
|
814
|
-
"required": ["city"],
|
|
815
|
-
},
|
|
816
|
-
"function": "tests.test_chat_generator.get_weather",
|
|
817
|
-
"outputs_to_string": None,
|
|
818
|
-
"inputs_from_state": None,
|
|
819
|
-
"outputs_to_state": None,
|
|
820
|
-
},
|
|
821
|
-
}
|
|
822
|
-
]
|
|
823
|
-
},
|
|
824
|
-
}
|
|
825
|
-
assert data["init_parameters"]["tools"] == expected_tools_data
|
|
830
|
+
# deserializing the serialized component must reproduce the original toolset
|
|
831
|
+
loaded = TransformersChatGenerator.from_dict(data)
|
|
832
|
+
assert isinstance(loaded.tools, Toolset)
|
|
833
|
+
assert list(loaded.tools) == list(toolset)
|
|
826
834
|
|
|
827
835
|
@pytest.mark.asyncio
|
|
828
836
|
async def test_run_async_with_streaming_callback(self, model_info_mock, mock_pipeline_with_tokenizer):
|
|
@@ -865,7 +873,7 @@ class TestTransformersChatGeneratorAsync:
|
|
|
865
873
|
@pytest.mark.integration
|
|
866
874
|
@pytest.mark.flaky(reruns=3, reruns_delay=10)
|
|
867
875
|
@pytest.mark.asyncio
|
|
868
|
-
async def test_live_run_async_with_streaming(self,
|
|
876
|
+
async def test_live_run_async_with_streaming(self, del_hf_env_vars_if_empty):
|
|
869
877
|
"""Test async streaming with a live model."""
|
|
870
878
|
streaming_chunks = []
|
|
871
879
|
|
|
@@ -873,7 +881,10 @@ class TestTransformersChatGeneratorAsync:
|
|
|
873
881
|
streaming_chunks.append(chunk)
|
|
874
882
|
|
|
875
883
|
llm = TransformersChatGenerator(
|
|
876
|
-
model="Qwen/Qwen3-0.6B",
|
|
884
|
+
model="Qwen/Qwen3-0.6B",
|
|
885
|
+
generation_kwargs={"max_new_tokens": 50},
|
|
886
|
+
streaming_callback=streaming_callback,
|
|
887
|
+
device=ComponentDevice.from_str("cpu"),
|
|
877
888
|
)
|
|
878
889
|
|
|
879
890
|
response = await llm.run_async(
|
|
@@ -640,7 +640,7 @@ def test_warm_up_use_hf_token(mocked_automodel, mocked_autotokenizer, initialize
|
|
|
640
640
|
"haystack_integrations.components.readers.transformers.extractive_reader."
|
|
641
641
|
"AutoModelForQuestionAnswering.from_pretrained"
|
|
642
642
|
)
|
|
643
|
-
def test_device_map_auto(mocked_automodel, _mocked_autotokenizer,
|
|
643
|
+
def test_device_map_auto(mocked_automodel, _mocked_autotokenizer, del_hf_env_vars_if_empty):
|
|
644
644
|
reader = TransformersExtractiveReader("deepset/roberta-base-squad2", model_kwargs={"device_map": "auto"})
|
|
645
645
|
auto_device = ComponentDevice.resolve_device(None)
|
|
646
646
|
|
|
@@ -651,7 +651,9 @@ def test_device_map_auto(mocked_automodel, _mocked_autotokenizer, del_hf_env_var
|
|
|
651
651
|
mocked_automodel.return_value = MockedModel()
|
|
652
652
|
reader.warm_up()
|
|
653
653
|
|
|
654
|
-
mocked_automodel.assert_called_once_with(
|
|
654
|
+
mocked_automodel.assert_called_once_with(
|
|
655
|
+
"deepset/roberta-base-squad2", token=reader.token.resolve_value(), device_map="auto"
|
|
656
|
+
)
|
|
655
657
|
assert reader.device == ComponentDevice.from_multiple(DeviceMap.from_hf({"": auto_device.to_hf()}))
|
|
656
658
|
|
|
657
659
|
|
|
@@ -660,7 +662,7 @@ def test_device_map_auto(mocked_automodel, _mocked_autotokenizer, del_hf_env_var
|
|
|
660
662
|
"haystack_integrations.components.readers.transformers.extractive_reader."
|
|
661
663
|
"AutoModelForQuestionAnswering.from_pretrained"
|
|
662
664
|
)
|
|
663
|
-
def test_device_map_str(mocked_automodel, _mocked_autotokenizer,
|
|
665
|
+
def test_device_map_str(mocked_automodel, _mocked_autotokenizer, del_hf_env_vars_if_empty):
|
|
664
666
|
reader = TransformersExtractiveReader("deepset/roberta-base-squad2", model_kwargs={"device_map": "cpu:0"})
|
|
665
667
|
|
|
666
668
|
class MockedModel:
|
|
@@ -670,7 +672,9 @@ def test_device_map_str(mocked_automodel, _mocked_autotokenizer, del_hf_env_vars
|
|
|
670
672
|
mocked_automodel.return_value = MockedModel()
|
|
671
673
|
reader.warm_up()
|
|
672
674
|
|
|
673
|
-
mocked_automodel.assert_called_once_with(
|
|
675
|
+
mocked_automodel.assert_called_once_with(
|
|
676
|
+
"deepset/roberta-base-squad2", token=reader.token.resolve_value(), device_map="cpu:0"
|
|
677
|
+
)
|
|
674
678
|
assert reader.device == ComponentDevice.from_multiple(DeviceMap.from_hf({"": "cpu:0"}))
|
|
675
679
|
|
|
676
680
|
|
|
@@ -679,7 +683,7 @@ def test_device_map_str(mocked_automodel, _mocked_autotokenizer, del_hf_env_vars
|
|
|
679
683
|
"haystack_integrations.components.readers.transformers.extractive_reader."
|
|
680
684
|
"AutoModelForQuestionAnswering.from_pretrained"
|
|
681
685
|
)
|
|
682
|
-
def test_device_map_dict(mocked_automodel, _mocked_autotokenizer,
|
|
686
|
+
def test_device_map_dict(mocked_automodel, _mocked_autotokenizer, del_hf_env_vars_if_empty):
|
|
683
687
|
reader = TransformersExtractiveReader(
|
|
684
688
|
"deepset/roberta-base-squad2", model_kwargs={"device_map": {"layer_1": 1, "classifier": "cpu"}}
|
|
685
689
|
)
|
|
@@ -692,7 +696,9 @@ def test_device_map_dict(mocked_automodel, _mocked_autotokenizer, del_hf_env_var
|
|
|
692
696
|
reader.warm_up()
|
|
693
697
|
|
|
694
698
|
mocked_automodel.assert_called_once_with(
|
|
695
|
-
"deepset/roberta-base-squad2",
|
|
699
|
+
"deepset/roberta-base-squad2",
|
|
700
|
+
token=reader.token.resolve_value(),
|
|
701
|
+
device_map={"layer_1": 1, "classifier": "cpu"},
|
|
696
702
|
)
|
|
697
703
|
assert reader.device == ComponentDevice.from_multiple(DeviceMap.from_hf({"layer_1": 1, "classifier": "cpu"}))
|
|
698
704
|
|
|
@@ -907,11 +913,9 @@ class TestDeduplication:
|
|
|
907
913
|
|
|
908
914
|
|
|
909
915
|
@pytest.mark.integration
|
|
910
|
-
def test_t5(
|
|
911
|
-
reader = TransformersExtractiveReader("sjrhuschlee/flan-t5-base-squad2")
|
|
912
|
-
answers = reader.run(example_queries[0], example_documents[0], top_k=2)[
|
|
913
|
-
"answers"
|
|
914
|
-
] # remove indices when batching support is reintroduced
|
|
916
|
+
def test_t5(del_hf_env_vars_if_empty):
|
|
917
|
+
reader = TransformersExtractiveReader("sjrhuschlee/flan-t5-base-squad2", device=ComponentDevice.from_str("cpu"))
|
|
918
|
+
answers = reader.run(example_queries[0], example_documents[0], top_k=2)["answers"]
|
|
915
919
|
assert answers[0].data == "Olaf Scholz"
|
|
916
920
|
assert answers[0].score == pytest.approx(0.8085031509399414, abs=1e-5)
|
|
917
921
|
assert answers[1].data == "Angela Merkel"
|
|
@@ -919,22 +923,12 @@ def test_t5(del_hf_env_vars):
|
|
|
919
923
|
assert answers[2].data is None
|
|
920
924
|
assert answers[2].score == pytest.approx(0.0378925803599941, abs=1e-5)
|
|
921
925
|
assert len(answers) == 3
|
|
922
|
-
# Uncomment assertions below when batching is reintroduced
|
|
923
|
-
# assert answers[0][2].score == pytest.approx(0.051331606147570596)
|
|
924
|
-
# assert answers[1][0].data == "Jerry"
|
|
925
|
-
# assert answers[1][0].score == pytest.approx(0.7413333654403687)
|
|
926
|
-
# assert answers[1][1].data == "Olaf Scholz"
|
|
927
|
-
# assert answers[1][1].score == pytest.approx(0.7266613841056824)
|
|
928
|
-
# assert answers[1][2].data is None
|
|
929
|
-
# assert answers[1][2].score == pytest.approx(0.0707035798685709)
|
|
930
926
|
|
|
931
927
|
|
|
932
928
|
@pytest.mark.integration
|
|
933
|
-
def test_roberta(
|
|
934
|
-
reader = TransformersExtractiveReader("deepset/tinyroberta-squad2")
|
|
935
|
-
answers = reader.run(example_queries[0], example_documents[0], top_k=2)[
|
|
936
|
-
"answers"
|
|
937
|
-
] # remove indices when batching is reintroduced
|
|
929
|
+
def test_roberta(del_hf_env_vars_if_empty):
|
|
930
|
+
reader = TransformersExtractiveReader("deepset/tinyroberta-squad2", device=ComponentDevice.from_str("cpu"))
|
|
931
|
+
answers = reader.run(example_queries[0], example_documents[0], top_k=2)["answers"]
|
|
938
932
|
assert answers[0].data == "Olaf Scholz"
|
|
939
933
|
assert answers[0].score == pytest.approx(0.8614975214004517)
|
|
940
934
|
assert answers[1].data == "Angela Merkel"
|
|
@@ -942,16 +936,3 @@ def test_roberta(del_hf_env_vars):
|
|
|
942
936
|
assert answers[2].data is None
|
|
943
937
|
assert answers[2].score == pytest.approx(0.019673851661650588, abs=1e-5)
|
|
944
938
|
assert len(answers) == 3
|
|
945
|
-
# uncomment assertions below when there is batching in v2
|
|
946
|
-
# assert answers[0][0].data == "Olaf Scholz"
|
|
947
|
-
# assert answers[0][0].score == pytest.approx(0.8614975214004517)
|
|
948
|
-
# assert answers[0][1].data == "Angela Merkel"
|
|
949
|
-
# assert answers[0][1].score == pytest.approx(0.857952892780304)
|
|
950
|
-
# assert answers[0][2].data is None
|
|
951
|
-
# assert answers[0][2].score == pytest.approx(0.0196738764278237)
|
|
952
|
-
# assert answers[1][0].data == "Jerry"
|
|
953
|
-
# assert answers[1][0].score == pytest.approx(0.7048940658569336)
|
|
954
|
-
# assert answers[1][1].data == "Olaf Scholz"
|
|
955
|
-
# assert answers[1][1].score == pytest.approx(0.6604189872741699)
|
|
956
|
-
# assert answers[1][2].data is None
|
|
957
|
-
# assert answers[1][2].score == pytest.approx(0.1002123719777046)
|
{transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/tests/test_named_entity_extractor.py
RENAMED
|
@@ -101,7 +101,7 @@ def test_named_entity_extractor_serde():
|
|
|
101
101
|
_ = TransformersNamedEntityExtractor.from_dict(serde_data)
|
|
102
102
|
|
|
103
103
|
|
|
104
|
-
def test_to_dict_default(
|
|
104
|
+
def test_to_dict_default(del_hf_env_vars_if_empty):
|
|
105
105
|
component = TransformersNamedEntityExtractor(
|
|
106
106
|
model="dslim/bert-base-NER",
|
|
107
107
|
device=ComponentDevice.from_str("mps"),
|
|
@@ -144,7 +144,7 @@ def test_to_dict_with_parameters():
|
|
|
144
144
|
}
|
|
145
145
|
|
|
146
146
|
|
|
147
|
-
def test_named_entity_extractor_from_dict_no_default_parameters(
|
|
147
|
+
def test_named_entity_extractor_from_dict_no_default_parameters(del_hf_env_vars_if_empty):
|
|
148
148
|
data = {
|
|
149
149
|
"type": COMPONENT_TYPE,
|
|
150
150
|
"init_parameters": {"model": "dslim/bert-base-NER"},
|
|
@@ -226,16 +226,16 @@ def test_named_entity_extractor_run_fails_with_wrong_number_of_annotations():
|
|
|
226
226
|
|
|
227
227
|
|
|
228
228
|
@pytest.mark.integration
|
|
229
|
-
def test_ner_extractor_init(
|
|
230
|
-
extractor = TransformersNamedEntityExtractor(model="dslim/bert-base-NER")
|
|
229
|
+
def test_ner_extractor_init(del_hf_env_vars_if_empty):
|
|
230
|
+
extractor = TransformersNamedEntityExtractor(model="dslim/bert-base-NER", device=ComponentDevice.from_str("cpu"))
|
|
231
231
|
extractor.warm_up()
|
|
232
232
|
assert extractor.initialized
|
|
233
233
|
|
|
234
234
|
|
|
235
235
|
@pytest.mark.integration
|
|
236
236
|
@pytest.mark.parametrize("batch_size", [1, 3])
|
|
237
|
-
def test_ner_extractor(raw_texts, hf_annotations, batch_size,
|
|
238
|
-
extractor = TransformersNamedEntityExtractor(model="dslim/bert-base-NER")
|
|
237
|
+
def test_ner_extractor(raw_texts, hf_annotations, batch_size, del_hf_env_vars_if_empty):
|
|
238
|
+
extractor = TransformersNamedEntityExtractor(model="dslim/bert-base-NER", device=ComponentDevice.from_str("cpu"))
|
|
239
239
|
extractor.warm_up()
|
|
240
240
|
|
|
241
241
|
_extract_and_check_predictions(extractor, raw_texts, hf_annotations, batch_size)
|
|
@@ -248,7 +248,7 @@ def test_ner_extractor(raw_texts, hf_annotations, batch_size, del_hf_env_vars):
|
|
|
248
248
|
reason="Export an env var called HF_API_TOKEN or HF_TOKEN containing the Hugging Face token to run this test.",
|
|
249
249
|
)
|
|
250
250
|
def test_ner_extractor_private_models(raw_texts, hf_annotations, batch_size):
|
|
251
|
-
extractor = TransformersNamedEntityExtractor(model="deepset/bert-base-NER")
|
|
251
|
+
extractor = TransformersNamedEntityExtractor(model="deepset/bert-base-NER", device=ComponentDevice.from_str("cpu"))
|
|
252
252
|
extractor.warm_up()
|
|
253
253
|
|
|
254
254
|
_extract_and_check_predictions(extractor, raw_texts, hf_annotations, batch_size)
|
|
@@ -256,11 +256,11 @@ def test_ner_extractor_private_models(raw_texts, hf_annotations, batch_size):
|
|
|
256
256
|
|
|
257
257
|
@pytest.mark.integration
|
|
258
258
|
@pytest.mark.parametrize("batch_size", [1, 3])
|
|
259
|
-
def test_ner_extractor_in_pipeline(raw_texts, hf_annotations, batch_size,
|
|
259
|
+
def test_ner_extractor_in_pipeline(raw_texts, hf_annotations, batch_size, del_hf_env_vars_if_empty):
|
|
260
260
|
pipeline = Pipeline()
|
|
261
261
|
pipeline.add_component(
|
|
262
262
|
name="ner_extractor",
|
|
263
|
-
instance=TransformersNamedEntityExtractor(model="dslim/bert-base-NER"),
|
|
263
|
+
instance=TransformersNamedEntityExtractor(model="dslim/bert-base-NER", device=ComponentDevice.from_str("cpu")),
|
|
264
264
|
)
|
|
265
265
|
|
|
266
266
|
outputs = pipeline.run(
|
|
@@ -54,7 +54,7 @@ class TestTransformersTextRouter:
|
|
|
54
54
|
}
|
|
55
55
|
|
|
56
56
|
@patch("haystack_integrations.components.routers.transformers.text_router.AutoConfig.from_pretrained")
|
|
57
|
-
def test_from_dict(self, mock_auto_config_from_pretrained,
|
|
57
|
+
def test_from_dict(self, mock_auto_config_from_pretrained, del_hf_env_vars_if_empty):
|
|
58
58
|
mock_auto_config_from_pretrained.return_value = MagicMock(label2id={"en": 0, "de": 1})
|
|
59
59
|
data = {
|
|
60
60
|
"type": COMPONENT_TYPE,
|
|
@@ -79,11 +79,11 @@ class TestTransformersTextRouter:
|
|
|
79
79
|
"model": "papluca/xlm-roberta-base-language-detection",
|
|
80
80
|
"device": ComponentDevice.resolve_device(None).to_hf(),
|
|
81
81
|
"task": "text-classification",
|
|
82
|
-
"token":
|
|
82
|
+
"token": component.token.resolve_value(),
|
|
83
83
|
}
|
|
84
84
|
|
|
85
85
|
@patch("haystack_integrations.components.routers.transformers.text_router.AutoConfig.from_pretrained")
|
|
86
|
-
def test_from_dict_no_default_parameters(self, mock_auto_config_from_pretrained,
|
|
86
|
+
def test_from_dict_no_default_parameters(self, mock_auto_config_from_pretrained, del_hf_env_vars_if_empty):
|
|
87
87
|
mock_auto_config_from_pretrained.return_value = MagicMock(label2id={"en": 0, "de": 1})
|
|
88
88
|
data = {
|
|
89
89
|
"type": COMPONENT_TYPE,
|
|
@@ -99,11 +99,11 @@ class TestTransformersTextRouter:
|
|
|
99
99
|
"model": "papluca/xlm-roberta-base-language-detection",
|
|
100
100
|
"device": ComponentDevice.resolve_device(None).to_hf(),
|
|
101
101
|
"task": "text-classification",
|
|
102
|
-
"token":
|
|
102
|
+
"token": component.token.resolve_value(),
|
|
103
103
|
}
|
|
104
104
|
|
|
105
105
|
@patch("haystack_integrations.components.routers.transformers.text_router.AutoConfig.from_pretrained")
|
|
106
|
-
def test_from_dict_with_cpu_device(self, mock_auto_config_from_pretrained,
|
|
106
|
+
def test_from_dict_with_cpu_device(self, mock_auto_config_from_pretrained, del_hf_env_vars_if_empty):
|
|
107
107
|
mock_auto_config_from_pretrained.return_value = MagicMock(label2id={"en": 0, "de": 1})
|
|
108
108
|
data = {
|
|
109
109
|
"type": COMPONENT_TYPE,
|
|
@@ -128,7 +128,7 @@ class TestTransformersTextRouter:
|
|
|
128
128
|
"model": "papluca/xlm-roberta-base-language-detection",
|
|
129
129
|
"device": ComponentDevice.from_str("cpu").to_hf(),
|
|
130
130
|
"task": "text-classification",
|
|
131
|
-
"token":
|
|
131
|
+
"token": component.token.resolve_value(),
|
|
132
132
|
}
|
|
133
133
|
|
|
134
134
|
@patch("haystack_integrations.components.routers.transformers.text_router.AutoConfig.from_pretrained")
|
|
@@ -172,8 +172,10 @@ class TestTransformersTextRouter:
|
|
|
172
172
|
assert out == {"en": "What is the color of the sky?"}
|
|
173
173
|
|
|
174
174
|
@pytest.mark.integration
|
|
175
|
-
def test_run(self,
|
|
176
|
-
router = TransformersTextRouter(
|
|
175
|
+
def test_run(self, del_hf_env_vars_if_empty):
|
|
176
|
+
router = TransformersTextRouter(
|
|
177
|
+
model="papluca/xlm-roberta-base-language-detection", device=ComponentDevice.from_str("cpu")
|
|
178
|
+
)
|
|
177
179
|
out = router.run("What is the color of the sky?")
|
|
178
180
|
assert set(router.labels) == {
|
|
179
181
|
"ar",
|
|
@@ -201,7 +203,11 @@ class TestTransformersTextRouter:
|
|
|
201
203
|
assert out == {"en": "What is the color of the sky?"}
|
|
202
204
|
|
|
203
205
|
@pytest.mark.integration
|
|
204
|
-
def test_wrong_labels(self,
|
|
205
|
-
router = TransformersTextRouter(
|
|
206
|
+
def test_wrong_labels(self, del_hf_env_vars_if_empty):
|
|
207
|
+
router = TransformersTextRouter(
|
|
208
|
+
model="papluca/xlm-roberta-base-language-detection",
|
|
209
|
+
labels=["en", "de"],
|
|
210
|
+
device=ComponentDevice.from_str("cpu"),
|
|
211
|
+
)
|
|
206
212
|
with pytest.raises(ValueError):
|
|
207
213
|
router.warm_up()
|
|
@@ -10,7 +10,7 @@ import torch
|
|
|
10
10
|
from haystack.utils.device import ComponentDevice
|
|
11
11
|
from transformers import AutoTokenizer, PreTrainedTokenizerFast
|
|
12
12
|
|
|
13
|
-
from haystack_integrations.
|
|
13
|
+
from haystack_integrations.common.transformers.utils import (
|
|
14
14
|
_resolve_hf_device_map,
|
|
15
15
|
_StopWordsCriteria,
|
|
16
16
|
)
|
|
@@ -52,7 +52,7 @@ class TestTransformersZeroShotDocumentClassifier:
|
|
|
52
52
|
},
|
|
53
53
|
}
|
|
54
54
|
|
|
55
|
-
def test_from_dict(self,
|
|
55
|
+
def test_from_dict(self, del_hf_env_vars_if_empty):
|
|
56
56
|
data = {
|
|
57
57
|
"type": COMPONENT_TYPE,
|
|
58
58
|
"init_parameters": {
|
|
@@ -80,10 +80,10 @@ class TestTransformersZeroShotDocumentClassifier:
|
|
|
80
80
|
"model": "cross-encoder/nli-deberta-v3-xsmall",
|
|
81
81
|
"device": ComponentDevice.resolve_device(None).to_hf(),
|
|
82
82
|
"task": "zero-shot-classification",
|
|
83
|
-
"token":
|
|
83
|
+
"token": component.token.resolve_value(),
|
|
84
84
|
}
|
|
85
85
|
|
|
86
|
-
def test_from_dict_no_default_parameters(self,
|
|
86
|
+
def test_from_dict_no_default_parameters(self, del_hf_env_vars_if_empty):
|
|
87
87
|
data = {
|
|
88
88
|
"type": COMPONENT_TYPE,
|
|
89
89
|
"init_parameters": {"model": "cross-encoder/nli-deberta-v3-xsmall", "labels": ["positive", "negative"]},
|
|
@@ -98,7 +98,7 @@ class TestTransformersZeroShotDocumentClassifier:
|
|
|
98
98
|
"model": "cross-encoder/nli-deberta-v3-xsmall",
|
|
99
99
|
"device": ComponentDevice.resolve_device(None).to_hf(),
|
|
100
100
|
"task": "zero-shot-classification",
|
|
101
|
-
"token":
|
|
101
|
+
"token": component.token.resolve_value(),
|
|
102
102
|
}
|
|
103
103
|
|
|
104
104
|
@patch("haystack_integrations.components.classifiers.transformers.zero_shot_document_classifier.pipeline")
|
|
@@ -165,9 +165,11 @@ class TestTransformersZeroShotDocumentClassifier:
|
|
|
165
165
|
component.run(documents=documents)
|
|
166
166
|
|
|
167
167
|
@pytest.mark.integration
|
|
168
|
-
def test_run(self,
|
|
168
|
+
def test_run(self, del_hf_env_vars_if_empty):
|
|
169
169
|
component = TransformersZeroShotDocumentClassifier(
|
|
170
|
-
model="cross-encoder/nli-deberta-v3-xsmall",
|
|
170
|
+
model="cross-encoder/nli-deberta-v3-xsmall",
|
|
171
|
+
labels=["positive", "negative"],
|
|
172
|
+
device=ComponentDevice.from_str("cpu"),
|
|
171
173
|
)
|
|
172
174
|
positive_document = Document(content="That's good. I like it. " * 1000)
|
|
173
175
|
negative_document = Document(content="That's bad. I don't like it.")
|
{transformers_haystack-0.1.0 → transformers_haystack-0.3.0}/tests/test_zero_shot_text_router.py
RENAMED
|
@@ -22,6 +22,7 @@ class TestTransformersZeroShotTextRouter:
|
|
|
22
22
|
"type": COMPONENT_TYPE,
|
|
23
23
|
"init_parameters": {
|
|
24
24
|
"labels": ["query", "passage"],
|
|
25
|
+
"multi_label": False,
|
|
25
26
|
"token": {"env_vars": ["HF_API_TOKEN", "HF_TOKEN"], "strict": False, "type": "env_var"},
|
|
26
27
|
"huggingface_pipeline_kwargs": {
|
|
27
28
|
"model": "MoritzLaurer/deberta-v3-base-zeroshot-v1.1-all-33",
|
|
@@ -31,7 +32,15 @@ class TestTransformersZeroShotTextRouter:
|
|
|
31
32
|
},
|
|
32
33
|
}
|
|
33
34
|
|
|
34
|
-
def
|
|
35
|
+
def test_multi_label_survives_a_serialization_round_trip(self, del_hf_env_vars_if_empty):
|
|
36
|
+
"""`multi_label` changes how the pipeline normalizes scores, so losing it changes routing decisions."""
|
|
37
|
+
router = TransformersZeroShotTextRouter(labels=["query", "passage"], multi_label=True)
|
|
38
|
+
|
|
39
|
+
restored = TransformersZeroShotTextRouter.from_dict(router.to_dict())
|
|
40
|
+
|
|
41
|
+
assert restored.multi_label is True
|
|
42
|
+
|
|
43
|
+
def test_from_dict(self, del_hf_env_vars_if_empty):
|
|
35
44
|
data = {
|
|
36
45
|
"type": COMPONENT_TYPE,
|
|
37
46
|
"init_parameters": {
|
|
@@ -55,10 +64,10 @@ class TestTransformersZeroShotTextRouter:
|
|
|
55
64
|
"model": "MoritzLaurer/deberta-v3-base-zeroshot-v1.1-all-33",
|
|
56
65
|
"device": ComponentDevice.resolve_device(None).to_hf(),
|
|
57
66
|
"task": "zero-shot-classification",
|
|
58
|
-
"token":
|
|
67
|
+
"token": component.token.resolve_value(),
|
|
59
68
|
}
|
|
60
69
|
|
|
61
|
-
def test_from_dict_no_default_parameters(self,
|
|
70
|
+
def test_from_dict_no_default_parameters(self, del_hf_env_vars_if_empty):
|
|
62
71
|
data = {
|
|
63
72
|
"type": COMPONENT_TYPE,
|
|
64
73
|
"init_parameters": {"labels": ["query", "passage"]},
|
|
@@ -73,7 +82,7 @@ class TestTransformersZeroShotTextRouter:
|
|
|
73
82
|
"model": "MoritzLaurer/deberta-v3-base-zeroshot-v1.1-all-33",
|
|
74
83
|
"device": ComponentDevice.resolve_device(None).to_hf(),
|
|
75
84
|
"task": "zero-shot-classification",
|
|
76
|
-
"token":
|
|
85
|
+
"token": component.token.resolve_value(),
|
|
77
86
|
}
|
|
78
87
|
|
|
79
88
|
@patch("haystack_integrations.components.routers.transformers.zero_shot_text_router.pipeline")
|
|
@@ -110,8 +119,8 @@ class TestTransformersZeroShotTextRouter:
|
|
|
110
119
|
assert out == {"query": "What is the color of the sky?"}
|
|
111
120
|
|
|
112
121
|
@pytest.mark.integration
|
|
113
|
-
def test_run(self,
|
|
114
|
-
router = TransformersZeroShotTextRouter(labels=["query", "passage"])
|
|
122
|
+
def test_run(self, del_hf_env_vars_if_empty):
|
|
123
|
+
router = TransformersZeroShotTextRouter(labels=["query", "passage"], device=ComponentDevice.from_str("cpu"))
|
|
115
124
|
out = router.run("What is the color of the sky?")
|
|
116
125
|
assert router.pipeline is not None
|
|
117
126
|
assert out == {"query": "What is the color of the sky?"}
|
|
@@ -1,26 +0,0 @@
|
|
|
1
|
-
# SPDX-FileCopyrightText: 2022-present deepset GmbH <info@deepset.ai>
|
|
2
|
-
#
|
|
3
|
-
# SPDX-License-Identifier: Apache-2.0
|
|
4
|
-
|
|
5
|
-
import pytest
|
|
6
|
-
from haystack.document_stores.in_memory import InMemoryDocumentStore
|
|
7
|
-
|
|
8
|
-
|
|
9
|
-
@pytest.fixture()
|
|
10
|
-
def del_hf_env_vars(monkeypatch):
|
|
11
|
-
"""
|
|
12
|
-
Delete Hugging Face environment variables for tests.
|
|
13
|
-
|
|
14
|
-
Prevents passing empty tokens to Hugging Face, which would cause API calls to fail.
|
|
15
|
-
This is particularly relevant for PRs opened from forks, where secrets are not available
|
|
16
|
-
and empty environment variables might be set instead of being removed.
|
|
17
|
-
|
|
18
|
-
See https://github.com/deepset-ai/haystack/issues/8811 for more details.
|
|
19
|
-
"""
|
|
20
|
-
monkeypatch.delenv("HF_API_TOKEN", raising=False)
|
|
21
|
-
monkeypatch.delenv("HF_TOKEN", raising=False)
|
|
22
|
-
|
|
23
|
-
|
|
24
|
-
@pytest.fixture()
|
|
25
|
-
def in_memory_doc_store():
|
|
26
|
-
return InMemoryDocumentStore()
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|