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:
@@ -52,7 +52,7 @@ def test_all_device():
|
||||
assert DistributedInterface().get_local_world_size() == int(os.getenv("LOCAL_WORLD_SIZE", "1"))
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cuda", "npu"])
|
||||
@pytest.mark.runs_on(["cuda", "npu", "xpu"])
|
||||
@pytest.mark.require_distributed(2)
|
||||
def test_multi_device():
|
||||
master_port = find_available_port()
|
||||
|
||||
@@ -79,6 +79,8 @@ def _get_visible_devices_env() -> str | None:
|
||||
return "CUDA_VISIBLE_DEVICES"
|
||||
elif CURRENT_DEVICE == "npu":
|
||||
return "ASCEND_RT_VISIBLE_DEVICES"
|
||||
elif CURRENT_DEVICE == "xpu":
|
||||
return "ZE_AFFINITY_MASK"
|
||||
else:
|
||||
return None
|
||||
|
||||
@@ -172,6 +174,8 @@ def _manage_distributed_env(request: FixtureRequest, monkeypatch: MonkeyPatch) -
|
||||
monkeypatch.setattr(torch.cuda, "device_count", lambda: 1)
|
||||
elif CURRENT_DEVICE == "npu":
|
||||
monkeypatch.setattr(torch.npu, "device_count", lambda: 1)
|
||||
elif CURRENT_DEVICE == "xpu":
|
||||
monkeypatch.setattr(torch.xpu, "device_count", lambda: 1)
|
||||
|
||||
|
||||
@pytest.fixture(scope="session", autouse=True)
|
||||
|
||||
@@ -96,7 +96,7 @@ def _test_sequence_parallel_loss(
|
||||
assert loss is not None
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cuda", "npu"])
|
||||
@pytest.mark.runs_on(["cuda", "npu", "xpu"])
|
||||
@pytest.mark.require_distributed(2)
|
||||
@pytest.mark.parametrize(("cp_size", "dp_size", "batch_size"), [(2, 1, 1), (2, 1, 2)])
|
||||
def test_sequence_parallel_loss(cp_size, dp_size, batch_size):
|
||||
|
||||
@@ -19,7 +19,7 @@ from llamafactory.v1.core.model_engine import ModelEngine
|
||||
from llamafactory.v1.samplers.cli_sampler import SyncSampler
|
||||
|
||||
|
||||
@pytest.mark.runs_on(["cuda", "npu"])
|
||||
@pytest.mark.runs_on(["cuda", "npu", "xpu"])
|
||||
def test_sync_sampler():
|
||||
model_args = ModelArguments(model="Qwen/Qwen3-4B-Instruct-2507")
|
||||
sample_args = SampleArguments()
|
||||
|
||||
@@ -21,7 +21,7 @@ import pytest
|
||||
|
||||
|
||||
@pytest.mark.xfail(reason="CI machines may OOM when heavily loaded.")
|
||||
@pytest.mark.runs_on(["cuda", "npu"])
|
||||
@pytest.mark.runs_on(["cuda", "npu", "xpu"])
|
||||
def test_fsdp2_dpo_trainer(tmp_path: Path):
|
||||
"""Test FSDP2 DPO trainer with sigmoid loss by simulating `llamafactory-cli dpo config.yaml`."""
|
||||
config_yaml = """\
|
||||
|
||||
@@ -20,7 +20,7 @@ import pytest
|
||||
|
||||
|
||||
@pytest.mark.xfail(reason="CI machines may OOM when heavily loaded.")
|
||||
@pytest.mark.runs_on(["cuda", "npu"])
|
||||
@pytest.mark.runs_on(["cuda", "npu", "xpu"])
|
||||
def test_fsdp2_sft_trainer(tmp_path: Path):
|
||||
"""Test FSDP2 SFT trainer by simulating `llamafactory-cli sft config.yaml` behavior."""
|
||||
config_yaml = """\
|
||||
|
||||
Reference in New Issue
Block a user