mirror of
https://github.com/hiyouga/LLaMA-Factory.git
synced 2026-10-05 14:25:43 +08:00
[xpu] add Intel XPU support to test infrastructure and runs_on markers (#10835)
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
@@ -37,13 +37,13 @@ MESSAGES = [
|
||||
EXPECTED_RESPONSE = "_rho"
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_chat():
|
||||
chat_model = ChatModel(INFER_ARGS)
|
||||
assert chat_model.chat(MESSAGES)[0].response_text == EXPECTED_RESPONSE
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_stream_chat():
|
||||
chat_model = ChatModel(INFER_ARGS)
|
||||
response = ""
|
||||
|
||||
@@ -49,7 +49,7 @@ INFER_ARGS = {
|
||||
OS_NAME = os.getenv("OS_NAME", "")
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
@pytest.mark.parametrize(
|
||||
"stage,dataset",
|
||||
[
|
||||
@@ -66,7 +66,7 @@ def test_run_exp(stage: str, dataset: str):
|
||||
assert os.path.exists(output_dir)
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_export():
|
||||
export_dir = os.path.join("output", "llama3_export")
|
||||
export_model({"export_dir": export_dir, **INFER_ARGS})
|
||||
|
||||
Reference in New Issue
Block a user