Skip to content

Commit 1531004

Browse files
fix image token output (#4487)
* fix * fix * fix * add test case * add test case * add test case
1 parent 31ac6c6 commit 1531004

3 files changed

Lines changed: 100 additions & 0 deletions

File tree

fastdeploy/engine/common_engine.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -718,6 +718,8 @@ def _zmq_send_generated_tokens(self):
718718
content.outputs.token_ids = token_ids
719719
content.outputs.text = delta_text
720720
new_contents.append(content)
721+
elif content.finished:
722+
new_contents.append(content)
721723
else:
722724
llm_logger.warning(
723725
f"current tokens need to accumulate, req_id: {request_id} {content.outputs.token_ids}"

tests/engine/test_decode_token.py

Lines changed: 51 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,51 @@
1+
import unittest
2+
from unittest.mock import MagicMock, patch
3+
4+
5+
class DummyDataProcessor:
6+
def __init__(self):
7+
self.decode_status = {}
8+
9+
def ids2tokens(self, token_ids, req_id):
10+
return "", [], None
11+
12+
13+
class TestDecodeToken(unittest.TestCase):
14+
@patch("fastdeploy.engine.common_engine.EngineSevice.__init__", return_value=None)
15+
def setUp(self, mock_init):
16+
from fastdeploy.engine.common_engine import EngineSevice
17+
18+
self.obj = EngineSevice(None)
19+
self.obj.data_processor = DummyDataProcessor()
20+
21+
@patch("fastdeploy.engine.common_engine.envs.FD_ENABLE_RETURN_TEXT", True)
22+
def test_decode_token_with_text(self):
23+
"""测试:env 启用 + 返回非空 delta_text"""
24+
self.obj.data_processor.ids2tokens = MagicMock(return_value=("hello", [10, 11, 12, 13], None))
25+
self.obj.data_processor.decode_status = {"req_1": (1, 3)}
26+
27+
delta_text, token_ids = self.obj._decode_token([1, 2, 3], "req_1", is_end=False)
28+
29+
assert delta_text == "hello"
30+
assert token_ids == [11, 12]
31+
32+
@patch("fastdeploy.engine.common_engine.envs.FD_ENABLE_RETURN_TEXT", True)
33+
def test_decode_token_empty_text(self):
34+
"""测试:env 启用 + 返回空 delta_text"""
35+
self.obj.data_processor.ids2tokens = MagicMock(return_value=("", [10, 11, 12], None))
36+
self.obj.data_processor.decode_status = {"req_1": (0, 2)}
37+
38+
delta_text, token_ids = self.obj._decode_token([1, 2], "req_1", is_end=False)
39+
40+
assert delta_text == ""
41+
assert token_ids == []
42+
43+
@patch("fastdeploy.engine.common_engine.envs.FD_ENABLE_RETURN_TEXT", True)
44+
def test_decode_token_with_is_end(self):
45+
"""测试:is_end=True 时 decode_status 被删除"""
46+
self.obj.data_processor.ids2tokens = MagicMock(return_value=("bye", [1, 2, 3, 4], None))
47+
self.obj.data_processor.decode_status = {"req_2": (0, 2)}
48+
49+
delta_text, token_ids = self.obj._decode_token([1, 2, 3], "req_2", is_end=True)
50+
51+
assert "req_2" not in self.obj.data_processor.decode_status

tests/engine/test_send_tokens.py

Lines changed: 47 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,47 @@
1+
import time
2+
from unittest import TestCase
3+
from unittest.mock import MagicMock, patch
4+
5+
6+
class TestZmqSendGeneratedTokens(TestCase):
7+
@patch("time.sleep", return_value=None)
8+
@patch("fastdeploy.engine.common_engine.EngineSevice.__init__", return_value=None)
9+
def setUp(self, mock_init, mock_sleep):
10+
from fastdeploy.engine.common_engine import EngineSevice
11+
12+
self.obj = EngineSevice(None)
13+
self.obj.running = True
14+
15+
# mock 依赖组件
16+
self.obj.scheduler = MagicMock()
17+
self.obj.send_response_server = MagicMock()
18+
self.obj._decode_token = MagicMock()
19+
self.obj._decode_token.return_value = ("decoded_text", [101, 102])
20+
self.obj.llm_logger = MagicMock()
21+
22+
def test_zmq_send_generated_tokens_normal_case(self):
23+
mock_output = MagicMock()
24+
mock_output.outputs.decode_type = 0
25+
mock_output.outputs.token_ids = [1, 2, 3]
26+
mock_output.finished = True
27+
28+
self.obj.scheduler.get_results.side_effect = [
29+
{"req_1": [mock_output]},
30+
{},
31+
]
32+
33+
def stop_running():
34+
time.sleep(0.01)
35+
self.obj.running = False
36+
37+
import threading
38+
39+
threading.Thread(target=stop_running).start()
40+
41+
self.obj._zmq_send_generated_tokens()
42+
43+
self.obj.send_response_server.send_response.assert_called_once()
44+
args, kwargs = self.obj.send_response_server.send_response.call_args
45+
assert args[0] == "req_1"
46+
assert isinstance(args[1], list)
47+
assert args[1][0].outputs.text == "decoded_text"

0 commit comments

Comments
 (0)