[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:
shubham singhal
2026-09-28 09:17:33 +05:30
committed by GitHub
parent 97b32d3133
commit d4823eb10e
21 changed files with 114 additions and 106 deletions

View File

@@ -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",
[