提交

[Misc] Black auto fix. (#4691)

Co-authored-by: Steve <ubuntu@ip-172-31-34-29.ap-northeast-1.compute.internal>
这个提交包含在:
Hongzhi (Steve), Chen
2022-10-10 10:53:51 +08:00
提交者 GitHub
父节点 c24e285a8c
当前提交 98325b1097
修改 48 个文件,包含 3247 行新增1757 行删除
+2 -2
查看文件
@@ -1,10 +1,10 @@
"""Feature storage classes for DataLoading"""
from .. import backend as F
from .base import *
from .numpy import *
# Defines the name TensorStorage
if F.get_preferred_backend() == 'pytorch':
if F.get_preferred_backend() == "pytorch":
from .pytorch_tensor import PyTorchTensorStorage as TensorStorage
else:
from .tensor import BaseTensorStorage as TensorStorage
+19 -9
查看文件
@@ -2,16 +2,19 @@
import threading
STORAGE_WRAPPERS = {}
def register_storage_wrapper(type_):
"""Decorator that associates a type to a ``FeatureStorage`` object.
"""
"""Decorator that associates a type to a ``FeatureStorage`` object."""
def deco(cls):
STORAGE_WRAPPERS[type_] = cls
return cls
return deco
def wrap_storage(storage):
"""Wrap an object into a FeatureStorage as specified by the ``register_storage_wrapper``
decorators.
@@ -20,11 +23,14 @@ def wrap_storage(storage):
if isinstance(storage, type_):
return storage_cls(storage)
assert isinstance(storage, FeatureStorage), (
"The frame column must be a tensor or a FeatureStorage object, got {}"
.format(type(storage)))
assert isinstance(
storage, FeatureStorage
), "The frame column must be a tensor or a FeatureStorage object, got {}".format(
type(storage)
)
return storage
class _FuncWrapper(object):
def __init__(self, func):
self.func = func
@@ -32,18 +38,21 @@ class _FuncWrapper(object):
def __call__(self, buf, *args):
buf[0] = self.func(*args)
class ThreadedFuture(object):
"""Wraps a function into a future asynchronously executed by a Python
``threading.Thread`. The function is being executed upon instantiation of
this object.
"""
def __init__(self, target, args):
self.buf = [None]
thread = threading.Thread(
target=_FuncWrapper(target),
args=[self.buf] + list(args),
daemon=True)
daemon=True,
)
thread.start()
self.thread = thread
@@ -52,14 +61,15 @@ class ThreadedFuture(object):
self.thread.join()
return self.buf[0]
class FeatureStorage(object):
"""Feature storage object which should support a fetch() operation. It is the
counterpart of a tensor for homogeneous graphs, or a dict of tensor for heterogeneous
graphs where the keys are node/edge types.
"""
def requires_ddp(self):
"""Whether the FeatureStorage requires the DataLoader to set use_ddp.
"""
"""Whether the FeatureStorage requires the DataLoader to set use_ddp."""
return False
def fetch(self, indices, device, pin_memory=False, **kwargs):
+7 -2
查看文件
@@ -1,11 +1,14 @@
"""Feature storage for ``numpy.memmap`` object."""
import numpy as np
from .base import FeatureStorage, ThreadedFuture, register_storage_wrapper
from .. import backend as F
from .base import FeatureStorage, ThreadedFuture, register_storage_wrapper
@register_storage_wrapper(np.memmap)
class NumpyStorage(FeatureStorage):
"""FeatureStorage that asynchronously reads features from a ``numpy.memmap`` object."""
def __init__(self, arr):
self.arr = arr
@@ -17,4 +20,6 @@ class NumpyStorage(FeatureStorage):
# pylint: disable=unused-argument
def fetch(self, indices, device, pin_memory=False, **kwargs):
return ThreadedFuture(target=self._fetch, args=(indices, device, pin_memory))
return ThreadedFuture(
target=self._fetch, args=(indices, device, pin_memory)
)
+27 -12
查看文件
@@ -1,43 +1,58 @@
"""Feature storages for PyTorch tensors."""
import torch
from ..utils import gather_pinned_tensor_rows
from .base import register_storage_wrapper
from .tensor import BaseTensorStorage
from ..utils import gather_pinned_tensor_rows
def _fetch_cpu(indices, tensor, feature_shape, device, pin_memory, **kwargs):
result = torch.empty(
indices.shape[0], *feature_shape, dtype=tensor.dtype,
pin_memory=pin_memory)
indices.shape[0],
*feature_shape,
dtype=tensor.dtype,
pin_memory=pin_memory,
)
torch.index_select(tensor, 0, indices, out=result)
kwargs['non_blocking'] = pin_memory
kwargs["non_blocking"] = pin_memory
result = result.to(device, **kwargs)
return result
def _fetch_cuda(indices, tensor, device, **kwargs):
return torch.index_select(tensor, 0, indices).to(device, **kwargs)
@register_storage_wrapper(torch.Tensor)
class PyTorchTensorStorage(BaseTensorStorage):
"""Feature storages for slicing a PyTorch tensor."""
def fetch(self, indices, device, pin_memory=False, **kwargs):
device = torch.device(device)
storage_device_type = self.storage.device.type
indices_device_type = indices.device.type
if storage_device_type != 'cuda':
if indices_device_type == 'cuda':
if storage_device_type != "cuda":
if indices_device_type == "cuda":
if self.storage.is_pinned():
return gather_pinned_tensor_rows(self.storage, indices)
else:
raise ValueError(
f'Got indices on device {indices.device} whereas the feature tensor '
f'is on {self.storage.device}. Please either (1) move the graph '
f'to GPU with to() method, or (2) pin the graph with '
f'pin_memory_() method.')
f"Got indices on device {indices.device} whereas the feature tensor "
f"is on {self.storage.device}. Please either (1) move the graph "
f"to GPU with to() method, or (2) pin the graph with "
f"pin_memory_() method."
)
# CPU to CPU or CUDA - use pin_memory and async transfer if possible
else:
return _fetch_cpu(indices, self.storage, self.storage.shape[1:], device,
pin_memory, **kwargs)
return _fetch_cpu(
indices,
self.storage,
self.storage.shape[1:],
device,
pin_memory,
**kwargs,
)
else:
# CUDA to CUDA or CPU
return _fetch_cuda(indices, self.storage, device, **kwargs)
+6 -2
查看文件
@@ -1,13 +1,17 @@
"""Feature storages for tensors across different frameworks."""
from .base import FeatureStorage
from .. import backend as F
from .base import FeatureStorage
class BaseTensorStorage(FeatureStorage):
"""FeatureStorage that synchronously slices features from a tensor and transfers
it to the given device.
"""
def __init__(self, tensor):
self.storage = tensor
def fetch(self, indices, device, pin_memory=False, **kwargs): # pylint: disable=unused-argument
def fetch(
self, indices, device, pin_memory=False, **kwargs
): # pylint: disable=unused-argument
return F.copy_to(F.gather_row(tensor, indices), device, **kwargs)