nemo_rl.distributed.named_sharding
nemo_rl.distributed.named_sharding
Contenu du Module
Classes
API
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]
Renvoie la forme de la disposition des rangs.
Renvoie les noms des axes.
Renvoie le nombre de dimensions.
Renvoie le nombre total de rangs.
Renvoie le tableau NumPy sous-jacent représentant la disposition.
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.
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.
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.
Obtient l’index numérique d’un axe nommé.
Obtient la taille d’un axe nommé.