99
1010
1111def 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
108114def 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 )
0 commit comments