Passer à la navigation

nemo_rl.distributed.collectives

Afficher en Markdown

Contenu du Module

Fonctions

NomDescription
rebalance_nd_tensorPrend des tenseurs avec des tailles de premier axe variables (à dim=0) et les empile en un seul tenseur.
gather_jagged_object_listsRassemble des listes de listes irrégulières d’objets sérialisables de tous les rangs et les aplatit en une seule liste.

Données

T

API

nemo_rl.distributed.collectives.T

Valeur: TypeVar(...)

nemo_rl.distributed.collectives.rebalance_nd_tensor(
tensor: torch.Tensor,
group: typing.Optional[torch.distributed.ProcessGroup] = None
) -> torch.Tensor

Prend des tenseurs avec des tailles de premier axe variables (à dim=0) et les empile en un seul tenseur.

Cette fonction gère le cas où différents GPU ont des tenseurs avec des tailles de batch différentes et les combine en un seul tenseur équilibré sur tous les rangs.

Par exemple, avec 3 GPU : GPU0 : tenseur de forme [3, D] GPU1 : tenseur de forme [5, D] GPU2 : tenseur de forme [2, D]

Après rééquilibrage : Tous les GPU auront le même tenseur de forme [10, D] (3+5+2=10)

REMARQUE : suppose que toutes les autres dimensions (non nulles) sont égales.

nemo_rl.distributed.collectives.gather_jagged_object_lists(
local_objects: list[nemo_rl.distributed.collectives.T],
group: typing.Optional[torch.distributed.ProcessGroup] = None
) -> list[nemo_rl.distributed.collectives.T]

Rassemble des listes de listes irrégulières d’objets sérialisables de tous les rangs et les aplatit en une seule liste.

Cette fonction gère le cas où différents GPU ont des listes de longueurs différentes et les combine en une seule liste contenant tous les objets de tous les rangs.

Par exemple, avec 3 GPU : GPU0 : [obj0, obj1] GPU1 : [obj2, obj3, obj4] GPU2 : [obj5]

Après le rassemblement : Tous les GPU auront : [obj0, obj1, obj2, obj3, obj4, obj5]

AVERTISSEMENT : synchrone

Arguments : local_objects : Liste des objets à rassembler du rang actuel group : Groupe de processus optionnel

Retourne : Liste aplatie de tous les objets de tous les rangs dans l’ordre [rang0, rang1, …]