mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-07-24 19:43:50 +00:00
Revert "Change memory format only for images"
This reverts commit 864b0d9d7f.
This commit is contained in:
parent
864b0d9d7f
commit
83f2068274
2 changed files with 7 additions and 6 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)}")
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue