Skip to content

Commit 6cb35d5

Browse files
Add Intel XPU device support (#1764)
* Add Intel XPU device support Signed-off-by: SjeYinTeoIntel <sje.yin.teo@intel.com> * fix: wrap device validation error message for flake8 Signed-off-by: SjeYinTeoIntel <sje.yin.teo@intel.com> --------- Signed-off-by: SjeYinTeoIntel <sje.yin.teo@intel.com>
1 parent ea68050 commit 6cb35d5

3 files changed

Lines changed: 81 additions & 19 deletions

File tree

‎modelscope/utils/constant.py‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -565,6 +565,7 @@ class Devices:
565565
"""device used for training and inference"""
566566
cpu = 'cpu'
567567
gpu = 'gpu'
568+
xpu = 'xpu'
568569

569570

570571
# Supported extensions for text datasets.

‎modelscope/utils/device.py‎

Lines changed: 30 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -9,22 +9,24 @@
99

1010

1111
def verify_device(device_name):
12-
""" Verify device is valid, device should be either cpu, cuda, gpu, cuda:X or gpu:X.
12+
""" Verify device is valid, device should be either cpu, cuda, gpu, xpu, cuda:X, gpu:X or xpu:X.
1313
1414
Args:
15-
device (str): device str, should be either cpu, cuda, gpu, gpu:X or cuda:X
16-
where X is the ordinal for gpu device.
15+
device (str): device str, should be either cpu, cuda, gpu, xpu, gpu:X, cuda:X or xpu:X
16+
where X is the ordinal for the device.
1717
1818
Return:
1919
device info (tuple): device_type and device_id, if device_id is not set, will use 0 as default.
2020
"""
21-
err_msg = 'device should be either cpu, cuda, gpu, gpu:X or cuda:X where X is the ordinal for gpu device.'
21+
err_msg = (
22+
'device should be either cpu, cuda, gpu, xpu, gpu:X, cuda:X or xpu:X '
23+
'where X is the ordinal for the device.')
2224
assert device_name is not None and device_name != '', err_msg
2325
device_name = device_name.lower()
2426
eles = device_name.split(':')
2527
assert len(eles) <= 2, err_msg
2628
assert device_name is not None
27-
assert eles[0] in ['cpu', 'cuda', 'gpu'], err_msg
29+
assert eles[0] in ['cpu', 'cuda', 'gpu', 'xpu'], err_msg
2830
device_type = eles[0]
2931
device_id = None
3032
if len(eles) > 1:
@@ -33,6 +35,8 @@ def verify_device(device_name):
3335
device_type = Devices.gpu
3436
if device_type == Devices.gpu and device_id is None:
3537
device_id = 0
38+
if device_type == Devices.xpu and device_id is None:
39+
device_id = 0
3640
return device_type, device_id
3741

3842

@@ -77,6 +81,12 @@ def device_placement(framework, device_name='gpu:0'):
7781
else:
7882
logger.debug(
7983
'pytorch: cuda is not available, using cpu instead.')
84+
elif device_type == Devices.xpu:
85+
if hasattr(torch, 'xpu') and torch.xpu.is_available():
86+
torch.xpu.set_device(f'xpu:{device_id}')
87+
else:
88+
logger.debug(
89+
'pytorch: xpu is not available, using cpu instead.')
8090
yield
8191
else:
8292
yield
@@ -86,23 +96,19 @@ def create_device(device_name):
8696
""" create torch device
8797
8898
Args:
89-
device_name (str): cpu, gpu, gpu:0, cuda:0 etc.
99+
device_name (str): cpu, gpu, gpu:0, cuda:0, xpu, xpu:0 etc.
90100
"""
91101
import torch
92102
device_type, device_id = verify_device(device_name)
93-
use_cuda = False
94103
if device_type == Devices.gpu:
95-
use_cuda = True
96-
if not torch.cuda.is_available():
97-
logger.info('cuda is not available, using cpu instead.')
98-
use_cuda = False
99-
100-
if use_cuda:
101-
device = torch.device(f'cuda:{device_id}')
102-
else:
103-
device = torch.device('cpu')
104-
105-
return device
104+
if torch.cuda.is_available():
105+
return torch.device(f'cuda:{device_id}')
106+
logger.info('cuda is not available, using cpu instead.')
107+
elif device_type == Devices.xpu:
108+
if hasattr(torch, 'xpu') and torch.xpu.is_available():
109+
return torch.device(f'xpu:{device_id}')
110+
logger.info('xpu is not available, using cpu instead.')
111+
return torch.device('cpu')
106112

107113

108114
def get_device():
@@ -114,6 +120,12 @@ def get_device():
114120
device_id = f"cuda:{os.environ['LOCAL_RANK']}"
115121
else:
116122
device_id = 'cuda:0'
123+
elif hasattr(torch, 'xpu') and torch.xpu.is_available():
124+
if dist.is_available() and dist.is_initialized(
125+
) and 'LOCAL_RANK' in os.environ:
126+
device_id = f"xpu:{os.environ['LOCAL_RANK']}"
127+
else:
128+
device_id = 'xpu:0'
117129
else:
118130
device_id = 'cpu'
119131
return torch.device(device_id)

‎tests/utils/test_device.py‎

Lines changed: 50 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@
1010

1111
from modelscope.utils.constant import Frameworks
1212
from modelscope.utils.device import (create_device, device_placement,
13-
verify_device)
13+
get_device, verify_device)
1414

1515

1616
class DeviceTest(unittest.TestCase):
@@ -44,6 +44,14 @@ def test_verify(self):
4444
self.assertEqual(device_name, 'gpu')
4545
self.assertTrue(device_id == 1)
4646

47+
device_name, device_id = verify_device('xpu')
48+
self.assertEqual(device_name, 'xpu')
49+
self.assertTrue(device_id == 0)
50+
51+
device_name, device_id = verify_device('xpu:1')
52+
self.assertEqual(device_name, 'xpu')
53+
self.assertTrue(device_id == 1)
54+
4755
with self.assertRaises(AssertionError):
4856
verify_device('xgu')
4957

@@ -80,6 +88,41 @@ def test_create_device_torch(self):
8088
self.assertTrue(device.type == target_device_type)
8189
self.assertTrue(device.index == target_device_index)
8290

91+
def test_get_device_cpu(self):
92+
if torch.cuda.is_available() or (hasattr(torch, 'xpu')
93+
and torch.xpu.is_available()):
94+
self.skipTest('accelerator present, cpu fallback not exercised')
95+
device = get_device()
96+
self.assertIsInstance(device, torch.device)
97+
self.assertEqual(device.type, 'cpu')
98+
99+
@unittest.skipUnless(
100+
hasattr(torch, 'xpu') and torch.xpu.is_available(), 'no xpu')
101+
def test_get_device_xpu(self):
102+
device = get_device()
103+
self.assertIsInstance(device, torch.device)
104+
self.assertEqual(device.type, 'xpu')
105+
106+
def test_create_device_xpu_fallback(self):
107+
if hasattr(torch, 'xpu') and torch.xpu.is_available():
108+
self.skipTest('xpu present, fallback not exercised')
109+
device = create_device('xpu')
110+
self.assertIsInstance(device, torch.device)
111+
self.assertEqual(device.type, 'cpu')
112+
113+
@unittest.skipUnless(
114+
hasattr(torch, 'xpu') and torch.xpu.is_available(), 'no xpu')
115+
def test_create_device_xpu(self):
116+
device = create_device('xpu')
117+
self.assertIsInstance(device, torch.device)
118+
self.assertEqual(device.type, 'xpu')
119+
self.assertEqual(device.index, 0)
120+
121+
device = create_device('xpu:0')
122+
self.assertIsInstance(device, torch.device)
123+
self.assertEqual(device.type, 'xpu')
124+
self.assertEqual(device.index, 0)
125+
83126
def test_device_placement_cpu(self):
84127
with device_placement(Frameworks.torch, 'cpu'):
85128
pass
@@ -102,6 +145,12 @@ def test_device_placement_torch_gpu(self):
102145
if torch.cuda.is_available():
103146
self.assertEqual(torch.cuda.current_device(), 0)
104147

148+
@unittest.skipUnless(
149+
hasattr(torch, 'xpu') and torch.xpu.is_available(), 'no xpu')
150+
def test_device_placement_torch_xpu(self):
151+
with device_placement(Frameworks.torch, 'xpu:0'):
152+
self.assertEqual(torch.xpu.current_device(), 0)
153+
105154

106155
if __name__ == '__main__':
107156
unittest.main()

0 commit comments

Comments
 (0)