Revert "Change memory format only for images"

This reverts commit 864b0d9d7f.
This commit is contained in:
Antonin Raffin 2022-12-17 12:41:28 +01:00
parent 864b0d9d7f
commit 83f2068274
No known key found for this signature in database
GPG key ID: B8B48F65CAD6232C
2 changed files with 7 additions and 6 deletions

View file

@ -131,10 +131,9 @@ class BaseBuffer(ABC):
(may be useful to avoid changing things be reference)
:return:
"""
kwargs = {"memory_format": th.channels_last} if len(array.shape) >= 4 else {}
if copy:
return th.tensor(array).to(self.device, non_blocking=True, **kwargs)
return th.as_tensor(array).to(self.device, non_blocking=True, **kwargs)
return th.tensor(array).to(self.device, memory_format=th.channels_last, non_blocking=True)
return th.as_tensor(array).to(self.device, memory_format=th.channels_last, non_blocking=True)
@staticmethod
def _normalize_obs(

View file

@ -460,11 +460,13 @@ def obs_as_tensor(
:param device: PyTorch device
:return: PyTorch tensor of the observation on a desired device.
"""
kwargs = {"memory_format": th.channels_last} if len(obs.shape) >= 4 else {}
if isinstance(obs, np.ndarray):
return th.as_tensor(obs).to(device, non_blocking=True, **kwargs)
return th.as_tensor(obs).to(device, memory_format=th.channels_last, non_blocking=True)
elif isinstance(obs, dict):
return {key: th.as_tensor(_obs).to(device, non_blocking=True, **kwargs) for (key, _obs) in obs.items()}
return {
key: th.as_tensor(_obs).to(device, memory_format=th.channels_last, non_blocking=True)
for (key, _obs) in obs.items()
}
else:
raise Exception(f"Unrecognized type of observation {type(obs)}")