Passer à la navigation

nemo_rl.distributed.named_sharding

Afficher en Markdown

Contenu du Module

Classes

NomDescription
NamedShardingReprésente un arrangement N-dimensionnel de rangs avec des axes nommés, facilitant le partitionnement, la réplication et la collection de données en fonction de ces axes.

API

class nemo_rl.distributed.named_sharding.NamedSharding(layout: typing.Sequence[typing.Any] | numpy.ndarray, names: list[str])

Représente un arrangement N-dimensionnel de rangs avec des axes nommés, facilitant le partitionnement, la réplication et la collection de données en fonction de ces axes.

Exemple : layout = [ [[0, 1, 2, 3], [4, 5, 6, 7]], ] names = [“dp”, “pp”, “tp”]

Ceci représente DP=1, PP=2, TP=4

sharding = NamedSharding(layout, names) print(sharding.shape) # Sortie : (1, 2, 4) print(sharding.names) # Sortie : [‘dp’, ‘pp’, ‘tp’] print(sharding.get_ranks(dp=0, pp=1)) # Sortie : [4, 5, 6, 7]

shape: dict[str, int]

Renvoie la forme de la disposition des rangs.

names: list[str]

Renvoie les noms des axes.

ndim: int

Renvoie le nombre de dimensions.

size: int

Renvoie le nombre total de rangs.

layout: numpy.ndarray[tuple[int, ...], numpy.dtype[numpy.int32]]

Renvoie le tableau NumPy sous-jacent représentant la disposition.

get_worker_coords(worker_id: int) -> dict[str, int]

Obtient les coordonnées d’un ID de travailleur spécifique dans la disposition de partitionnement.

Args : worker_id : L’ID entier du travailleur.

Renvoie : Un dictionnaire mappant les noms d’axes à leurs coordonnées entières pour le worker_id donné.

Lève : ValueError : Si le worker_id n’est pas trouvé dans la disposition.

get_ranks_by_coord(**coords: int) -> list[int]

Obtient tous les rangs correspondant aux coordonnées spécifiées pour les axes nommés.

Args : **coords : Arguments de mot-clé où la clé est le nom de l’axe (par exemple, “dp”, “tp”) et la valeur est la coordonnée entière le long de cet axe. Les axes non spécifiés correspondront à toutes les coordonnées le long de cet axe.

Renvoie : Une liste triée de rangs entiers uniques correspondant aux critères de coordonnées donnés. Renvoie une liste vide si aucun rang ne correspond.

Lève : ValueError : Si un nom d’axe non valide est fourni.

get_ranks(**kwargs: int) -> typing.Union[nemo_rl.distributed.named_sharding.NamedSharding, int]

Obtient les rangs correspondant à des indices spécifiques le long des axes nommés.

Args : **kwargs : Arguments de mot-clé où la clé est le nom de l’axe (par exemple, “dp”, “tp”) et la valeur est l’index le long de cet axe.

Renvoie : Une nouvelle instance de NamedSharding représentant le sous-ensemble des rangs. La forme du partitionnement renvoyé correspond aux axes non spécifiés dans les kwargs. Si tous les axes sont spécifiés, un entier est renvoyé.

Lève : ValueError : Si un nom d’axe non valide est fourni ou si un index est hors limites.

get_axis_index(name: str) -> int

Obtient l’index numérique d’un axe nommé.

get_axis_size(name: str) -> int

Obtient la taille d’un axe nommé.

__repr__() -> str
__eq__(other: object) -> bool