refactor: 更新inference 部分的实现
This commit is contained in:
@@ -1,11 +1,11 @@
|
||||
"""Shared fixtures for inference tests."""
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from astrai.inference.server import app
|
||||
from astrai.inference.server import app, _engine
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -30,13 +30,17 @@ def mock_model_param():
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_generator(mock_model_param):
|
||||
"""Mock the GeneratorFactory and its generators."""
|
||||
with patch("astrai.inference.server.GeneratorFactory") as MockFactory:
|
||||
mock_gen = MagicMock()
|
||||
mock_gen.generate.return_value = "mock response"
|
||||
MockFactory.create.return_value = mock_gen
|
||||
yield MockFactory, mock_gen
|
||||
def mock_engine():
|
||||
"""Create a mock InferenceEngine."""
|
||||
mock = MagicMock()
|
||||
mock.generate.return_value = "mock response"
|
||||
mock.get_stats.return_value = {
|
||||
"total_tasks": 0,
|
||||
"total_tokens": 0,
|
||||
"active_tasks": 0,
|
||||
"waiting_queue": 0,
|
||||
}
|
||||
return mock
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
|
||||
@@ -6,24 +6,29 @@ import pytest
|
||||
def test_health_no_model(client, monkeypatch):
|
||||
"""GET /health should return 200 even when model not loaded."""
|
||||
monkeypatch.setattr("astrai.inference.server._model_param", None)
|
||||
monkeypatch.setattr("astrai.inference.server._engine", None)
|
||||
response = client.get("/health")
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["status"] == "ok"
|
||||
assert not data["model_loaded"]
|
||||
assert not data["engine_ready"]
|
||||
|
||||
|
||||
def test_health_with_model(client, loaded_model):
|
||||
def test_health_with_model(client, loaded_model, mock_engine, monkeypatch):
|
||||
"""GET /health should return 200 when model is loaded."""
|
||||
monkeypatch.setattr("astrai.inference.server._engine", mock_engine)
|
||||
response = client.get("/health")
|
||||
assert response.status_code == 200
|
||||
assert response.json() == {"status": "ok", "model_loaded": True}
|
||||
data = response.json()
|
||||
assert data["status"] == "ok"
|
||||
assert data["model_loaded"] is True
|
||||
assert data["engine_ready"] is True
|
||||
|
||||
|
||||
def test_generate_non_stream(client, loaded_model, mock_generator):
|
||||
def test_generate_non_stream(client, loaded_model, mock_engine, monkeypatch):
|
||||
"""POST /generate with stream=false should return JSON response."""
|
||||
MockFactory, mock_gen = mock_generator
|
||||
mock_gen.generate.return_value = "Test response"
|
||||
monkeypatch.setattr("astrai.inference.server._engine", mock_engine)
|
||||
response = client.post(
|
||||
"/generate",
|
||||
params={
|
||||
@@ -37,15 +42,19 @@ def test_generate_non_stream(client, loaded_model, mock_generator):
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["response"] == "Test response"
|
||||
MockFactory.create.assert_called_once()
|
||||
assert data["response"] == "mock response"
|
||||
|
||||
|
||||
def test_generate_stream(client, loaded_model, mock_generator):
|
||||
def test_generate_stream(client, loaded_model, mock_engine, monkeypatch):
|
||||
"""POST /generate with stream=true should return plain text stream."""
|
||||
MockFactory, mock_gen = mock_generator
|
||||
# Simulate a streaming generator that yields two chunks
|
||||
mock_gen.generate.return_value = ["chunk1", "chunk2"]
|
||||
|
||||
# Create a streaming mock
|
||||
def stream_gen():
|
||||
yield "chunk1"
|
||||
yield "chunk2"
|
||||
|
||||
mock_engine.generate.return_value = stream_gen()
|
||||
monkeypatch.setattr("astrai.inference.server._engine", mock_engine)
|
||||
response = client.post(
|
||||
"/generate",
|
||||
params={
|
||||
@@ -66,10 +75,10 @@ def test_generate_stream(client, loaded_model, mock_generator):
|
||||
assert "chunk2" in content
|
||||
|
||||
|
||||
def test_chat_completions_non_stream(client, loaded_model, mock_generator):
|
||||
def test_chat_completions_non_stream(client, loaded_model, mock_engine, monkeypatch):
|
||||
"""POST /v1/chat/completions with stream=false returns OpenAI‑style JSON."""
|
||||
MockFactory, mock_gen = mock_generator
|
||||
mock_gen.generate.return_value = "Assistant reply"
|
||||
mock_engine.generate.return_value = "Assistant reply"
|
||||
monkeypatch.setattr("astrai.inference.server._engine", mock_engine)
|
||||
response = client.post(
|
||||
"/v1/chat/completions",
|
||||
json={
|
||||
@@ -88,11 +97,17 @@ def test_chat_completions_non_stream(client, loaded_model, mock_generator):
|
||||
assert data["choices"][0]["message"]["content"] == "Assistant reply"
|
||||
|
||||
|
||||
def test_chat_completions_stream(client, loaded_model, mock_generator):
|
||||
def test_chat_completions_stream(client, loaded_model, mock_engine, monkeypatch):
|
||||
"""POST /v1/chat/completions with stream=true returns SSE stream."""
|
||||
MockFactory, mock_gen = mock_generator
|
||||
|
||||
# Simulate a streaming generator that yields cumulative responses
|
||||
mock_gen.generate.return_value = ["cumulative1", "cumulative2"]
|
||||
def stream_gen():
|
||||
yield "cumulative1"
|
||||
yield "cumulative2"
|
||||
yield "[DONE]"
|
||||
|
||||
mock_engine.generate.return_value = stream_gen()
|
||||
monkeypatch.setattr("astrai.inference.server._engine", mock_engine)
|
||||
response = client.post(
|
||||
"/v1/chat/completions",
|
||||
json={
|
||||
@@ -116,10 +131,9 @@ def test_chat_completions_stream(client, loaded_model, mock_generator):
|
||||
assert any("cumulative2" in line for line in lines)
|
||||
|
||||
|
||||
def test_generate_with_history(client, loaded_model, mock_generator):
|
||||
def test_generate_with_history(client, loaded_model, mock_engine, monkeypatch):
|
||||
"""POST /generate with history parameter."""
|
||||
MockFactory, mock_gen = mock_generator
|
||||
mock_gen.generate.return_value = "Response with history"
|
||||
monkeypatch.setattr("astrai.inference.server._engine", mock_engine)
|
||||
response = client.post(
|
||||
"/generate",
|
||||
params={
|
||||
@@ -129,12 +143,8 @@ def test_generate_with_history(client, loaded_model, mock_generator):
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
MockFactory.create.assert_called_once()
|
||||
# Check that history was passed correctly (currently history is not parsed due to FastAPI limitation)
|
||||
call_args = MockFactory.create.call_args
|
||||
req = call_args[0][1] # second argument is GenerationRequest
|
||||
# Because history cannot be passed via query params, it will be None
|
||||
assert req.history is None
|
||||
# Verify the engine.generate was called
|
||||
mock_engine.generate.assert_called_once()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -3,8 +3,6 @@ import os
|
||||
import torch
|
||||
|
||||
from astrai.config.param_config import ModelParameter
|
||||
from astrai.inference.generator import EmbeddingEncoderCore, GeneratorCore
|
||||
|
||||
|
||||
def test_model_parameter(test_env):
|
||||
save_dir = os.path.join(test_env["test_dir"], "save")
|
||||
@@ -32,40 +30,4 @@ def test_transformer(test_env):
|
||||
test_env["transformer_config"].max_len,
|
||||
test_env["transformer_config"].vocab_size,
|
||||
)
|
||||
assert output_logits.shape == target_shape
|
||||
|
||||
|
||||
# generator
|
||||
def test_embedding_encoder_core(test_env):
|
||||
parameter = ModelParameter(
|
||||
test_env["model"], test_env["tokenizer"], test_env["transformer_config"]
|
||||
)
|
||||
encoder = EmbeddingEncoderCore(parameter)
|
||||
|
||||
single_emb = encoder.encode("测试文本")
|
||||
assert isinstance(single_emb, torch.Tensor)
|
||||
assert single_emb.shape[-1] == test_env["transformer_config"].dim
|
||||
|
||||
batch_emb = encoder.encode(["测试1", "测试2"])
|
||||
assert isinstance(batch_emb, list)
|
||||
assert len(batch_emb) == 2
|
||||
|
||||
|
||||
def test_generator_core(test_env):
|
||||
parameter = ModelParameter(
|
||||
test_env["model"], test_env["tokenizer"], test_env["transformer_config"]
|
||||
)
|
||||
generator = GeneratorCore(parameter)
|
||||
input_ids = torch.randint(0, test_env["transformer_config"].vocab_size, (4, 10))
|
||||
next_token_id, cache_increase = generator.generate_iterator(
|
||||
input_ids=input_ids,
|
||||
temperature=0.8,
|
||||
top_k=50,
|
||||
top_p=0.95,
|
||||
attn_mask=None,
|
||||
kv_caches=None,
|
||||
start_pos=0,
|
||||
)
|
||||
|
||||
assert next_token_id.shape == (4, 1)
|
||||
assert cache_increase == 10
|
||||
assert output_logits.shape == target_shape
|
||||
Reference in New Issue
Block a user