refactor(processor): rename internal device variable for clarity (#1857)

- Changed the internal device variable from `_device` to `tensor_device` for improved readability and consistency.
- Updated references throughout the class to reflect the new variable name.
This commit is contained in:
Adil Zouitine
2025-09-03 18:39:06 +02:00
committed by GitHub
parent 75dcfd4886
commit f125d5e3bf
+4 -4
View File
@@ -50,8 +50,8 @@ class DeviceProcessorStep(ProcessorStep):
} }
def __post_init__(self): def __post_init__(self):
self._device: torch.device = get_safe_torch_device(self.device) self.tensor_device: torch.device = get_safe_torch_device(self.device)
self.device = self._device.type # cuda might have changed to cuda:1 self.device = self.tensor_device.type # cuda might have changed to cuda:1
self.non_blocking = "cuda" in str(self.device) self.non_blocking = "cuda" in str(self.device)
# Validate and convert float_dtype string to torch dtype # Validate and convert float_dtype string to torch dtype
@@ -73,7 +73,7 @@ class DeviceProcessorStep(ProcessorStep):
Otherwise, it moves to the configured device. Otherwise, it moves to the configured device.
""" """
# Determine target device # Determine target device
if tensor.is_cuda and self._device.type == "cuda": if tensor.is_cuda and self.tensor_device.type == "cuda":
# Both tensor and target are on GPU - preserve tensor's GPU placement # Both tensor and target are on GPU - preserve tensor's GPU placement
# This handles multi-GPU scenarios where Accelerate has already placed # This handles multi-GPU scenarios where Accelerate has already placed
# tensors on the correct GPU for each process # tensors on the correct GPU for each process
@@ -81,7 +81,7 @@ class DeviceProcessorStep(ProcessorStep):
else: else:
# Either tensor is on CPU, or we're configured for CPU # Either tensor is on CPU, or we're configured for CPU
# In both cases, use the configured device # In both cases, use the configured device
target_device = self._device target_device = self.tensor_device
# Only move if necessary # Only move if necessary
if tensor.device != target_device: if tensor.device != target_device: