/
usr
/
local
/
lib64
/
python3.6
/
site-packages
/
torch
/
distributed
/
_sharding_spec
/
/usr/local/lib64/python3.6/site-packages/torch/distributed/_sharding_spec
mkdir
upload
Name
Size
Mode
Actions
__pycache__/
-
0755
rm
api.py
3788
0644
edit
dl
rm
_internals.py
5711
0644
edit
dl
rm
__init__.py
153
0644
edit
dl
rm
Edit:
/usr/local/lib64/python3.6/site-packages/torch/distributed/_sharding_spec/api.py
(3788B)
from abc import ABC from dataclasses import dataclass from typing import List, Union import torch from ._internals import ( ShardMetadata, validate_non_overlapping_shards_metadata ) class PlacementSpec(ABC): """ Base class representing the placement of an entity. Subclasses of this class can be used to specify customized placements which might not be covered by existing APIs. """ pass @dataclass class DevicePlacementSpec(PlacementSpec): """ Associates placement of an entity with a single device. Args: device(:class:`torch.distributed._remote_device`): The device to place the entity on. """ device: torch.distributed._remote_device def __post_init__(self): if not isinstance(self.device, torch.distributed._remote_device): self.device = torch.distributed._remote_device(self.device) class ShardingSpec(PlacementSpec): """ Base class representing sharding specifications. It is special type of PlacementSpec. """ pass @dataclass class ChunkShardingSpec(ShardingSpec): """ This is a type of PlacementSpec that defines the placement as being sharded across multiple devices. In particular, it represents sharding a Tensor along a single dimension into equal chunks (similar to :meth:`torch.chunk`). The semantics of how a tensor is partitioned is inline with :meth:`torch.chunk`, where ``dim`` in torch.chunk corresponds to the specified ``dim`` and ``chunks`` in torch.chunk is the number of elements in the placement specified. Args: dim (int or str): The dimension to shard on, could be an integer representing the dimension or a string in case of named tensors where dimensions are named. placement(List[Union[_remote_device, str]]): Specifies the placement of each shard of the Tensor. The size of the list represents the number of shards to be created. This could be a list of :class:`torch.distributed._remote_device`'s. This list could also contain a string which represents remote device as accepted by :class:`torch.distributed._remote_device` """ ShardingDim = Union[int, str] dim: ShardingDim placements: List[Union[torch.distributed._remote_device, str]] def __post_init__(self): self._verify_dim(self.dim) for i, remote_device in enumerate(self.placements): if not isinstance(remote_device, torch.distributed._remote_device): self.placements[i] = torch.distributed._remote_device(remote_device) @staticmethod def _verify_dim(dim): if not (isinstance(dim, int) or isinstance(dim, str)): raise ValueError(f'{dim} needs to either be an int or str') @dataclass class EnumerableShardingSpec(ShardingSpec): """ This is a type of PlacementSpec that allows users to specify a generic sharding scheme by enumerating exactly how each shard is laid out. Args: shards(List[ShardMetadata]): List of :class:`ShardMetadata` objects representing each shard. Note that none of the shards should overlap. """ shards: List[ShardMetadata] def __post_init__(self): if len(self.shards) == 0: raise ValueError(f'Empty shard list provided: {self.shards}') # Validate each shard has same rank. rank = -1 for shard in self.shards: if rank != -1 and rank != len(shard.shard_offsets): raise ValueError(f'Found inconsistent ranks for shards: {rank} and {len(shard.shard_offsets)}') rank = len(shard.shard_offsets) validate_non_overlapping_shards_metadata(self.shards)
Save
cmd:
run