@@ -315,12 +315,31 @@ def get_cached_numpy(self, dtype, shape, generation_kind="input", scale=1.2):
315315 return get_cached_numpy_array (dtype , shape , generation_kind = generation_kind , scale = scale )
316316
317317 def _use_gpu (self , api_config = None , dtype = None ):
318+ """判断是否启用 GPU tensor 生成;不代表 Paddle kernel 的 place。"""
318319 if not is_gpu_mode ():
319320 return False
320321 if self .place is not None and "cpu" in str (self .place ).lower ():
321322 return False
322323 return "gpu" in paddle .device .get_device ()
323324
325+ def _paddle_kernel_uses_gpu (self , api_config ):
326+ """test_cpu 只决定 Paddle kernel,不能被 use_gpu_mode 覆盖。"""
327+ if self .place is not None and "cpu" in str (self .place ).lower ():
328+ return False
329+ if getattr (api_config , "test_cpu" , False ):
330+ return False
331+ return "gpu" in paddle .device .get_device ()
332+
333+ def _torch_source_device_for_paddle (self , api_config ):
334+ """返回送入 Paddle 前的 Torch 源设备,遵守 Paddle kernel place。"""
335+ if self .place is not None and "cpu" in str (self .place ).lower ():
336+ return torch .device ("cpu" )
337+ if getattr (api_config , "test_cpu" , False ):
338+ return torch .device ("cpu" )
339+ if self ._paddle_kernel_uses_gpu (api_config ):
340+ return torch .device ("cuda" , torch .cuda .current_device ())
341+ return torch .device ("cpu" )
342+
324343 def _supports_autograd (self , dtype = None ):
325344 dtype = dtype or self .dtype
326345 return dtype in AUTOGRAD_DTYPES
@@ -3433,10 +3452,12 @@ def get_paddle_tensor(self, api_config):
34333452 and self ._use_gpu (api_config )
34343453 ):
34353454 self .paddle_tensor = self ._make_gpu_paddle_tensor (api_config )
3455+ if getattr (api_config , "test_cpu" , False ):
3456+ self .paddle_tensor = self .paddle_tensor ._copy_to (paddle .CPUPlace (), False )
34363457 return self .paddle_tensor
34373458 if self .cpu_tensor is not None :
34383459 torch_tensor = self .cpu_tensor .to (
3439- device = torch . device ( "cuda:0" ) if self ._use_gpu (api_config ) else "cpu" ,
3460+ device = self ._torch_source_device_for_paddle (api_config ),
34403461 copy = True ,
34413462 )
34423463 self .paddle_tensor = paddle .utils .dlpack .from_dlpack (
@@ -3445,7 +3466,10 @@ def get_paddle_tensor(self, api_config):
34453466 self .paddle_tensor .stop_gradient = not self ._requires_autograd (api_config )
34463467 return self .paddle_tensor
34473468 if self .numpy_tensor is None and self ._use_gpu (api_config ):
3448- return self .get_gpu_paddle_tensor (api_config )
3469+ tensor = self .get_gpu_paddle_tensor (api_config )
3470+ if getattr (api_config , "test_cpu" , False ):
3471+ self .paddle_tensor = tensor ._copy_to (paddle .CPUPlace (), False )
3472+ return self .paddle_tensor
34493473 if not self .is_contiguous and self .strides is not None :
34503474 self .paddle_tensor = self ._create_strided_paddle_tensor (api_config )
34513475 print (
@@ -3462,10 +3486,11 @@ def get_paddle_tensor(self, api_config):
34623486 if self .dtype == "bfloat16"
34633487 else ("float16" if self .dtype in FLOAT8_DTYPES else self .dtype )
34643488 )
3489+ operator_place = paddle .CPUPlace () if getattr (api_config , "test_cpu" , False ) else self .place
34653490 self .paddle_tensor = paddle .to_tensor (
34663491 self .get_numpy_tensor (api_config ),
34673492 dtype = intermediate_dtype ,
3468- place = self . place ,
3493+ place = operator_place ,
34693494 )
34703495
34713496 if self .dtype == "bfloat16" :
@@ -3497,18 +3522,19 @@ def _create_strided_paddle_tensor(self, api_config):
34973522 try :
34983523 intermediate_dtype = "float16" if self .dtype in FLOAT8_DTYPES else self .dtype
34993524 storage_size = self ._strided_storage_size ()
3525+ operator_place = paddle .CPUPlace () if getattr (api_config , "test_cpu" , False ) else self .place
35003526 flat_tensor = paddle .zeros (
35013527 [storage_size ],
35023528 dtype = intermediate_dtype ,
3503- device = self . place ,
3529+ device = operator_place ,
35043530 )
35053531 tensor = paddle .as_strided (flat_tensor , self .shape , self .strides )
35063532 logical_tensor = self .get_numpy_tensor (api_config )
35073533 if logical_tensor .size > 0 :
35083534 tensor [...] = paddle .to_tensor (
35093535 logical_tensor ,
35103536 dtype = intermediate_dtype ,
3511- place = self . place ,
3537+ place = operator_place ,
35123538 )
35133539 if self .dtype in FLOAT8_DTYPES :
35143540 flat_tensor = paddle .cast (flat_tensor , dtype = self .dtype )
0 commit comments