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:
@@ -38,19 +38,19 @@ TOOLS = [
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_empty_formatter():
|
||||
formatter = EmptyFormatter(slots=["\n"])
|
||||
assert formatter.apply() == ["\n"]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_string_formatter():
|
||||
formatter = StringFormatter(slots=["<s>", "Human: {{content}}\nAssistant:"])
|
||||
assert formatter.apply(content="Hi") == ["<s>", "Human: Hi\nAssistant:"]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_function_formatter():
|
||||
formatter = FunctionFormatter(slots=["{{content}}", "</s>"], tool_format="default")
|
||||
tool_calls = json.dumps(FUNCTION)
|
||||
@@ -60,7 +60,7 @@ def test_function_formatter():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_multi_function_formatter():
|
||||
formatter = FunctionFormatter(slots=["{{content}}", "</s>"], tool_format="default")
|
||||
tool_calls = json.dumps([FUNCTION] * 2)
|
||||
@@ -71,7 +71,7 @@ def test_multi_function_formatter():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_default_tool_formatter():
|
||||
formatter = ToolFormatter(tool_format="default")
|
||||
assert formatter.apply(content=json.dumps(TOOLS)) == [
|
||||
@@ -90,14 +90,14 @@ def test_default_tool_formatter():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_default_tool_extractor():
|
||||
formatter = ToolFormatter(tool_format="default")
|
||||
result = """Action: test_tool\nAction Input: {"foo": "bar", "size": 10}"""
|
||||
assert formatter.extract(result) == [("test_tool", """{"foo": "bar", "size": 10}""")]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_default_multi_tool_extractor():
|
||||
formatter = ToolFormatter(tool_format="default")
|
||||
result = (
|
||||
@@ -110,14 +110,14 @@ def test_default_multi_tool_extractor():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_glm4_function_formatter():
|
||||
formatter = FunctionFormatter(slots=["{{content}}"], tool_format="glm4")
|
||||
tool_calls = json.dumps(FUNCTION)
|
||||
assert formatter.apply(content=tool_calls) == ["""tool_name\n{"foo": "bar", "size": 10}"""]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_glm4_tool_formatter():
|
||||
formatter = ToolFormatter(tool_format="glm4")
|
||||
assert formatter.apply(content=json.dumps(TOOLS)) == [
|
||||
@@ -128,14 +128,14 @@ def test_glm4_tool_formatter():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_glm4_tool_extractor():
|
||||
formatter = ToolFormatter(tool_format="glm4")
|
||||
result = """test_tool\n{"foo": "bar", "size": 10}\n"""
|
||||
assert formatter.extract(result) == [("test_tool", """{"foo": "bar", "size": 10}""")]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_llama3_function_formatter():
|
||||
formatter = FunctionFormatter(slots=["{{content}}<|eot_id|>"], tool_format="llama3")
|
||||
tool_calls = json.dumps(FUNCTION)
|
||||
@@ -144,7 +144,7 @@ def test_llama3_function_formatter():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_llama3_multi_function_formatter():
|
||||
formatter = FunctionFormatter(slots=["{{content}}<|eot_id|>"], tool_format="llama3")
|
||||
tool_calls = json.dumps([FUNCTION] * 2)
|
||||
@@ -155,7 +155,7 @@ def test_llama3_multi_function_formatter():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_llama3_tool_formatter():
|
||||
formatter = ToolFormatter(tool_format="llama3")
|
||||
date = datetime.now().strftime("%d %b %Y")
|
||||
@@ -169,14 +169,14 @@ def test_llama3_tool_formatter():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_llama3_tool_extractor():
|
||||
formatter = ToolFormatter(tool_format="llama3")
|
||||
result = """{"name": "test_tool", "parameters": {"foo": "bar", "size": 10}}\n"""
|
||||
assert formatter.extract(result) == [("test_tool", """{"foo": "bar", "size": 10}""")]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_llama3_multi_tool_extractor():
|
||||
formatter = ToolFormatter(tool_format="llama3")
|
||||
result = (
|
||||
@@ -189,7 +189,7 @@ def test_llama3_multi_tool_extractor():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_mistral_function_formatter():
|
||||
formatter = FunctionFormatter(slots=["[TOOL_CALLS] {{content}}", "</s>"], tool_format="mistral")
|
||||
tool_calls = json.dumps(FUNCTION)
|
||||
@@ -199,7 +199,7 @@ def test_mistral_function_formatter():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_mistral_multi_function_formatter():
|
||||
formatter = FunctionFormatter(slots=["[TOOL_CALLS] {{content}}", "</s>"], tool_format="mistral")
|
||||
tool_calls = json.dumps([FUNCTION] * 2)
|
||||
@@ -211,7 +211,7 @@ def test_mistral_multi_function_formatter():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_mistral_tool_formatter():
|
||||
formatter = ToolFormatter(tool_format="mistral")
|
||||
wrapped_tool = {"type": "function", "function": TOOLS[0]}
|
||||
@@ -220,14 +220,14 @@ def test_mistral_tool_formatter():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_mistral_tool_extractor():
|
||||
formatter = ToolFormatter(tool_format="mistral")
|
||||
result = """{"name": "test_tool", "arguments": {"foo": "bar", "size": 10}}"""
|
||||
assert formatter.extract(result) == [("test_tool", """{"foo": "bar", "size": 10}""")]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_mistral_multi_tool_extractor():
|
||||
formatter = ToolFormatter(tool_format="mistral")
|
||||
result = (
|
||||
@@ -240,7 +240,7 @@ def test_mistral_multi_tool_extractor():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_qwen_function_formatter():
|
||||
formatter = FunctionFormatter(slots=["{{content}}<|im_end|>\n"], tool_format="qwen")
|
||||
tool_calls = json.dumps(FUNCTION)
|
||||
@@ -249,7 +249,7 @@ def test_qwen_function_formatter():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_qwen_multi_function_formatter():
|
||||
formatter = FunctionFormatter(slots=["{{content}}<|im_end|>\n"], tool_format="qwen")
|
||||
tool_calls = json.dumps([FUNCTION] * 2)
|
||||
@@ -260,7 +260,7 @@ def test_qwen_multi_function_formatter():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_qwen_tool_formatter():
|
||||
formatter = ToolFormatter(tool_format="qwen")
|
||||
wrapped_tool = {"type": "function", "function": TOOLS[0]}
|
||||
@@ -274,14 +274,14 @@ def test_qwen_tool_formatter():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_qwen_tool_extractor():
|
||||
formatter = ToolFormatter(tool_format="qwen")
|
||||
result = """<tool_call>\n{"name": "test_tool", "arguments": {"foo": "bar", "size": 10}}\n</tool_call>"""
|
||||
assert formatter.extract(result) == [("test_tool", """{"foo": "bar", "size": 10}""")]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_qwen38_tool_formatter():
|
||||
formatter = ToolFormatter(tool_format="qwen3_8")
|
||||
wrapped_tool = {"type": "function", "function": TOOLS[0]}
|
||||
@@ -289,7 +289,7 @@ def test_qwen38_tool_formatter():
|
||||
assert json.dumps(wrapped_tool, ensure_ascii=False) in output
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_qwen_multi_tool_extractor():
|
||||
formatter = ToolFormatter(tool_format="qwen")
|
||||
result = (
|
||||
@@ -302,7 +302,7 @@ def test_qwen_multi_tool_extractor():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_lfm2_function_formatter():
|
||||
formatter = FunctionFormatter(slots=["{{content}}<|im_end|>\n"], tool_format="lfm2")
|
||||
tool_calls = json.dumps(FUNCTION)
|
||||
@@ -311,7 +311,7 @@ def test_lfm2_function_formatter():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_lfm2_multi_function_formatter():
|
||||
formatter = FunctionFormatter(slots=["{{content}}<|im_end|>\n"], tool_format="lfm2")
|
||||
tool_calls = json.dumps([FUNCTION] * 2)
|
||||
@@ -321,7 +321,7 @@ def test_lfm2_multi_function_formatter():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_lfm2_tool_formatter():
|
||||
formatter = ToolFormatter(tool_format="lfm2")
|
||||
assert formatter.apply(content=json.dumps(TOOLS)) == [
|
||||
@@ -329,14 +329,14 @@ def test_lfm2_tool_formatter():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_lfm2_tool_extractor():
|
||||
formatter = ToolFormatter(tool_format="lfm2")
|
||||
result = """<|tool_call_start|>[test_tool(foo="bar", size=10)]<|tool_call_end|>"""
|
||||
assert formatter.extract(result) == [("test_tool", """{"foo": "bar", "size": 10}""")]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_lfm2_multi_tool_extractor():
|
||||
formatter = ToolFormatter(tool_format="lfm2")
|
||||
result = """<|tool_call_start|>[test_tool(foo="bar", size=10), another_tool(foo="job", size=2)]<|tool_call_end|>"""
|
||||
@@ -346,7 +346,7 @@ def test_lfm2_multi_tool_extractor():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_lfm2_tool_extractor_with_nested_dict():
|
||||
formatter = ToolFormatter(tool_format="lfm2")
|
||||
result = """<|tool_call_start|>[search(query="test", options={"limit": 10, "offset": 0})]<|tool_call_end|>"""
|
||||
@@ -358,7 +358,7 @@ def test_lfm2_tool_extractor_with_nested_dict():
|
||||
assert args["options"] == {"limit": 10, "offset": 0}
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_lfm2_tool_extractor_with_list_arg():
|
||||
formatter = ToolFormatter(tool_format="lfm2")
|
||||
result = """<|tool_call_start|>[batch_process(items=[1, 2, 3], enabled=True)]<|tool_call_end|>"""
|
||||
@@ -370,7 +370,7 @@ def test_lfm2_tool_extractor_with_list_arg():
|
||||
assert args["enabled"] is True
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_lfm2_tool_extractor_no_match():
|
||||
formatter = ToolFormatter(tool_format="lfm2")
|
||||
result = "This is a regular response without tool calls."
|
||||
@@ -378,7 +378,7 @@ def test_lfm2_tool_extractor_no_match():
|
||||
assert extracted == result
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_lfm2_tool_round_trip():
|
||||
formatter = FunctionFormatter(slots=["{{content}}"], tool_format="lfm2")
|
||||
tool_formatter = ToolFormatter(tool_format="lfm2")
|
||||
@@ -390,7 +390,7 @@ def test_lfm2_tool_round_trip():
|
||||
assert json.loads(extracted[0][1]) == original["arguments"]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_minicpm5_function_formatter():
|
||||
formatter = FunctionFormatter(slots=["{{content}}<|im_end|>\n"], tool_format="minicpm5")
|
||||
tool_calls = json.dumps(FUNCTION)
|
||||
@@ -399,7 +399,7 @@ def test_minicpm5_function_formatter():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_minicpm5_multi_function_formatter():
|
||||
formatter = FunctionFormatter(slots=["{{content}}<|im_end|>\n"], tool_format="minicpm5")
|
||||
tool_calls = json.dumps([FUNCTION] * 2)
|
||||
@@ -411,7 +411,7 @@ def test_minicpm5_multi_function_formatter():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_minicpm5_tool_formatter():
|
||||
formatter = ToolFormatter(tool_format="minicpm5")
|
||||
wrapped = json.dumps({"type": "function", "function": TOOLS[0]}, ensure_ascii=False)
|
||||
@@ -427,28 +427,28 @@ def test_minicpm5_tool_formatter():
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_minicpm5_tool_extractor():
|
||||
formatter = ToolFormatter(tool_format="minicpm5")
|
||||
result = '<function name="test_tool"><param name="foo">bar</param><param name="size">10</param></function>'
|
||||
assert formatter.extract(result) == [("test_tool", """{"foo": "bar", "size": 10}""")]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_minicpm5_tool_extractor_cdata():
|
||||
formatter = ToolFormatter(tool_format="minicpm5")
|
||||
result = '<function name="test_tool"><param name="foo"><![CDATA[a < b\nsecond line]]></param></function>'
|
||||
assert formatter.extract(result) == [("test_tool", json.dumps({"foo": "a < b\nsecond line"}))]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
def test_minicpm5_tool_extractor_malformed_value():
|
||||
formatter = ToolFormatter(tool_format="minicpm5")
|
||||
result = '<function name="test_tool"><param name="foo">{[1, 2], [3, 4]}</param></function>'
|
||||
assert formatter.extract(result) == [("test_tool", json.dumps({"foo": "{[1, 2], [3, 4]}"}))]
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cpu", "mps"])
|
||||
@pytest.mark.runs_on(["cpu", "mps", "xpu"])
|
||||
@pytest.mark.parametrize(
|
||||
"arguments",
|
||||
[
|
||||
|
||||
Reference in New Issue
Block a user