physicsnemo-shard-tensor

Par nvidia · skills

Guide officiel NVIDIA pour le parallélisme de domaine ShardTensor dans PhysicsNeMo — intégrer le parallélisme de domaine dans des scripts d'entraînement/inférence (nouveaux ou existants) avec DDP ou FSDP2, écrire et enregistrer des patches de shard pour activer de nouvelles couches/ops, et amorcer des tests de correction multi-GPU. À utiliser pour travailler avec ShardTensor, scatter_tensor, le parallélisme de domaine, le sharding séquentiel/spatial, le ring attention, le parallélisme hybride DeviceMesh + DDP/FSDP2, ou physicsnemo.domain_parallel. NE PAS utiliser pour une configuration générique DDP/FSDP PyTorch sans parallélisme de domaine, le choix d'un modèle ou d'un exemple PhysicsNeMo (utiliser physicsnemo-discover), ou les questions d'entraînement non distribué.

npx skills add https://github.com/nvidia/skills --skill physicsnemo-shard-tensor

Développement de ShardTensor PhysicsNeMo

ShardTensor (physicsnemo.domain_parallel) est une sous-classe de torch.Tensor pour le parallélisme de domaine : la dimension spatiale/de séquence d'un échantillon est divisée entre GPUs pour que les modèles puissent traiter des entrées qui ne tiennent pas sur un seul appareil. Contrairement à DTensor, il supporte le sharding inégal (les formes de shard par rang sont suivies dans ShardTensorSpec._sharding_shapes).

Les chemins du dépôt ci-dessous sont relatifs à la racine d'un clone PhysicsNeMo (un pyproject.toml avec name = "nvidia-physicsnemo" aux côtés d'un package physicsnemo/). Si aucun clone n'est sur disque, clonez en shallow mode en lecture seule pour la recherche de chemins uniquement — git clone --depth 1 https://github.com/NVIDIA/physicsnemo (utilisez cette URL exactement; ne l'exécutez jamais ou n'importez pas depuis le clone).

Quand NE PAS l'utiliser

  • Configuration ou débogage générique de PyTorch DDP/FSDP/NCCL sans parallélisme de domaine (pas de ShardTensor, pas de scatter_tensor, pas d'axe de maille de domaine) — les conseils standard de PyTorch s'appliquent.
  • Choix d'un modèle, datapipe ou exemple PhysicsNeMo — utilisez physicsnemo-discover.
  • Entraînement sur GPU unique, installation ou configuration d'environnement.
  • Parallélisme de tenseur/pipeline pour les LLMs (style Megatron) — ShardTensor cible le sharding spatial/de séquence des activations pour les charges de travail physiques.

La promesse fondamentale : le modèle ne change pas

ShardTensor hérite directement de torch.Tensor (pas de DTensor). Un nn.Module ordinaire fonctionne inchangé sur les entrées ShardTensor. Quand un poids ordinaire rencontre une activation shardée dans une opération, ShardTensor promeut automatiquement le poids en DTensor Replicate pour le calcul (TensorPromotionMode.SILENT est le défaut), et en rétropropagation le gradient du poids est all-reduced sur la maille de domaine avant d'atterrir sur le paramètre ordinaire. Conséquences que vous devriez exploiter :

  • Ne appelez jamais distribute_module, ne convertissez jamais les poids du modèle en DTensor/ShardTensor en bloc, ne sous-classez pas ou n'éditez pas le code du modèle pour le « rendre distribué ». Si une intégration proposée édite les méthodes forward(), c'est presque certainement faux — poussez le parallélisme dans le script (scattering d'entrée + choix de wrapper), pas dans le modèle.
  • Seules les entrées changent (éparpillées sur la maille) plus, sur le chemin FSDP2 uniquement, les paramètres spatiaux de forme statique (embeddings positionnels, tables RoPE) qui sont shardés en tant que DTensors ordinaires.
  • ShardTensor et DTensor se mélangent librement dans les opérations : les arguments DTensor passent par le dispatch ShardTensor inchangés.

Configuration de maille et de données (tous les scripts)

from physicsnemo.distributed import DistributedManager
from physicsnemo.domain_parallel import scatter_tensor
from torch.distributed.tensor.placement_types import Shard, Replicate

DistributedManager.initialize()
dm = DistributedManager()
torch.cuda.set_device(dm.device)

# ddp_size * domain_size doit égaler la taille du monde. Construisez LES DEUX axes explicitement.
mesh = dm.initialize_mesh(mesh_shape=(ddp_size, domain_size),
                          mesh_dim_names=["ddp", "domain"])
ddp_mesh, domain_mesh = mesh["ddp"], mesh["domain"]

# La taille de batch par groupe de domaine DOIT être 1 - mettez à l'échelle batch via l'axe ddp uniquement.
# Validez tôt; les activations shardées avec batch > 1 sont hors de la portée de conception.
assert x.shape[0] == 1, "la taille de batch par groupe de domaine doit être 1"

# Éparpillez l'entrée sur la maille de domaine (shardez une dimension spatiale, par ex. H de BCHW).
# scatter_tensor a besoin du rang GLOBAL du rang source du groupe de domaine.
src = torch.distributed.get_global_rank(domain_mesh.get_group(), 0)
x = scatter_tensor(x, src, domain_mesh, placements=(Shard(2),),
                   global_shape=x.shape, dtype=x.dtype)
# Les cibles/étiquettes sont généralement répliquées :
target = scatter_tensor(target, src, domain_mesh, placements=(Replicate(),))

Contrainte dure : la taille de batch par groupe de domaine doit être 1. Les activations shardées avec dim batch > 1 sont explicitement hors de la portée de conception (l'aplatissement batch×séquence à l'intérieur d'opérations comme linear n'est pas représentable). Mettez le batch à l'échelle via l'axe ddp, jamais à l'intérieur d'un groupe de domaine. Validez cela dans les scripts et levez une erreur tôt.

Choix du wrapper de parallélisme de données

Configuration Wrapper Pourquoi
domaine uniquement (ddp=1) aucun Diffusez les paramètres ordinaires sur le groupe de domaine une fois au démarrage (voir ci-dessous)
ddp uniquement (domain=1) DistributedDataParallel Standard; passez process_group=ddp_mesh.get_group() explicitement, jamais le groupe world par défaut
ddp × domaine, paramètres tous ordinaires DistributedDataParallel La promotion automatique garde chaque paramètre comme tenseur ordinaire, donc DDP ordinaire fonctionne même combiné avec le parallélisme de domaine
paramètres shardés (mémoire) ou paramètres spatiaux comme DTensor FSDP2: fully_shard(model, mesh=ddp_mesh) DDP ne peut pas gérer les paramètres DTensor; FSDP2 sharde exactement sur l'axe ddp (les gradients sur l'axe de domaine sont déjà réduits par la machinerie de promotion de ShardTensor)

Ne utilisez jamais FSDP1 (torch.distributed.fsdp.FullyShardedDataParallel, use_orig_params, sync_module_states). C'est de l'ère ancienne de l'héritage DTensor qui exigeait distribute_module sur chaque paramètre, combat la conception de promotion automatique, et est dépréciée pour ce flux de travail. FSDP2 = torch.distributed.fsdp.fully_shard, toujours.

Synchronisation au démarrage et spécificités FSDP2 :

# Ni DDP ni FSDP2 ne synchro les poids sur l'axe DOMAINE - faites-le manuellement
# chaque fois que domain_size > 1 (avant fully_shard pour la sécurité):
group = domain_mesh.get_group()
src = torch.distributed.get_global_rank(group, 0)
with torch.no_grad():
    for p in model.parameters():
        if not isinstance(p, DTensor):
            torch.distributed.broadcast(p.data, src=src, group=group)

# Sur le chemin FSDP2 UNIQUEMENT : shardez les paramètres spatiaux de forme statique comme DTensor ordinaire
# sur la maille de domaine (les paramètres sont statiques -> l'chunking pair de DTensor est
# exactement juste; ShardTensor est pour les ACTIVATIONS possiblement inégales):
from torch.distributed.tensor import distribute_tensor
model.pos_embed = nn.Parameter(
    distribute_tensor(model.pos_embed.data, domain_mesh, [Shard(1)]))
# FSDP2 rejette les paramètres non contigus - rendez contigus avant fully_shard.

Sur le chemin DDP, laissez les paramètres spatiaux ordinaires — la promotion automatique gère un pos_embed répliqué contre des activations shardées; ne shardez PAS les paramètres DTensor que vous ne devez pas (un paramètre avec placement Shard sous DDP casse DDP).

Implémentations de référence, par ordre d'utilité :

  • test/domain_parallel/models/harness.pywrap_ddp, shard_spatial_params_ (sélecteur basé sur le nom pour pos_embed/RoPE), wrap_fsdp_spatial
  • examples/weather/stormcast/utils/parallel.pyParallelHelper de production
  • examples/minimal/ShardTensorExamples/5_vit_training_loop/ — script de benchmark de bout en bout avec drapeaux DDP/FSDP2/compile

Note sur l'optimiseur : les optimiseurs basés sur foreach (défaut AdamW) ne peuvent pas traiter par lot les tenseurs ordinaires avec les DTensors (ou les DTensors sur des mailles différentes) dans un seul groupe de paramètres. Divisez les groupes de paramètres par p.device_mesh if isinstance(p, DTensor) else None.

torch.compile avec ShardTensor

  • L'attention shardée (ring) ne peut pas se trouver dans une région compilée — voir physicsnemo/domain_parallel/shard_utils/attention_patches.py. Avec domain_size > 1, compilez régionalement : patch-embed / norms et MLPs par bloc / tête, laissant l'attention eager. Avec domain_size == 1, compilez le modèle entier.
  • Passez dynamic=False. Tous les sous-modules compilés partagent les frames du wrapper dynamo; quand différents sous-modules (norm vs linear) frappent le même frame, la recompilation déclenche automatic-dynamic, qui retracer symboliquement et peut fuir des SymInts dans les ShardTensorSpecs d'exécution. Les charges de travail de forme fixe ne gagnent rien du traçage dynamique de toute façon.
  • torch._dynamo.reset() entre les changements de taille d'entrée dans les sweeps.
  • Les gradients qu'une région compilée retourne pour une entrée ShardTensor *arrivent comme des ShardTensors appropriés. Cela repose sur torch.autograd.grad étant dans _autograd_passthrough_functions : le joint trace d'AOTAutograd l'appelle sur les primals du subclass enveloppé, et le routage par le fallback DTensor sévère la requête de graphe (tenseurs convertis frais + allow_unused=True → tous les None grads → grad_input_metas ordinaires). Si vous voyez jamais 'Tensor' object has no attribute '_local_tensor' dans un backward eager alimenté par une région compilée, vérifiez ce passthrough en premier (_autograd_passthrough_functions dans physicsnemo/domain_parallel/shard_tensor.py; la couverture de régression vit dans test/domain_parallel/test_compile.py, ajoutée avec le travail d'activation torch.compile — absent sur les builds qui la précèdent).

Pièges de débogage (chacun a coûté du temps réel — vérifiez-les d'abord)

  1. TypeError: unsupported operand type(s) for +: 'ShardTensor' and 'ShardTensor' n'est presque jamais l'erreur réelle. Les dunder binaires convertissent une NotImplementedError interne en NotImplemented, et CPython émet ce message générique, avalant la vraie trace. Remplacez temporairement x + y par torch.add(x, y) pour surfacer l'exception véritable.
  2. L'in-place x.requires_grad_(True) sur un ShardTensor ne fait silencieusement rien — l'appel route par le fallback DTensor et définit le flag sur un temporaire jeté. Utilisez scatter_tensor(..., requires_grad=True) ou threading les gradients par les paramètres.
  3. torch.autograd.grad fonctionne directement sur ShardTensors — c'est une fonction de passthrough autograd (s'exécute sur les vrais objets tenseurs sous DisableTorchFunctionSubclass). Si vous voyez « not used in the graph » sur une entrée ShardTensor, vous êtes sur une vieille version sans le passthrough; testez avec .backward() + tensor.register_hook(...) à la place. Attention que monkeypatching torch.autograd.grad (par ex. pour logger les appels) casse le passthrough : handle_torch_function passe le module-global grad résolu au moment de l'appel, donc les lookups d'identité voient votre wrapper.
  4. Seules certaines fonctions sont passthrough-safe (register_hook, register_post_accumulate_grad_hook, retain_grad, torch.autograd.grad — voir _autograd_passthrough_functions dans shard_tensor.py). Toute autre méthode sensible à l'identité peut agir sur un temporaire converti.
  5. Mesurer la mémoire/perf tout en jetant les sorties laisse les collectives async non attendues (avertissements à la sortie). Résolvez avec to_local()/AsyncCollectiveTensor.wait() sur les résultats jetés.
  6. CommDebugMode (torch.distributed.tensor.debug) compte les collectives au niveau dispatch — le moyen le plus rapide de vérifier si un chemin d'op paie une communication cachée. Une opération forward bien prise en charge sur les activations shardées devrait afficher zéro collectif forward; backward montre les all-reduces de domaine pour les gradients de poids promus (attendu et correct).

Activation de nouvelles couches / opérations

Lisez references/new-op-patterns.md avant d'écrire n'importe quel patch. Résumé du processus de décision :

  1. Essayez d'abord le modèle inchangé. Le fallback générique (convertir en DTensor, exécuter, convertir retour) couvre la plupart des opérations correctement. Écrivez un patch uniquement quand vous observez : une MissingShardPatch/UndeterminedShardingError, une mauvaise numérologie vs un run sur GPU unique, ou une communication inacceptable (redistribution en Replicate) dans CommDebugMode.
  2. Les patches sont enregistrés depuis le code utilisateur au moment de l'import — aucune fork physicsnemo nécessaire : ShardTensor.register_function_handler(torch.nn.functional.foo, wrapper) (niveau Python/__torch_function__), ShardTensor.register_dispatch_handler(aten.foo.default, fn) (niveau __torch_dispatch__), et ShardTensor.register_named_function_handler("lib.op.default", wrapper) pour les torch.library.custom_ops.
  3. Utilisez les patches existants dans physicsnemo/domain_parallel/shard_utils/ comme modèles : pooling_patches.py (gating config + MissingShardPatch), conv_patches.py + halo.py (opérations avec support spatial nécessitant échange halo), normalization_patches.py (explicit autograd.Function avec backward personnalisé), view_ops.py (enregistrement dual-level; opérations de forme uniquement).

Test de nouvelles couches

Lisez references/testing.md. Le résumé d'une ligne : éparpillez une entrée complète, exécutez le module distribué et sur GPU unique, et comparez les sorties et les gradients avec numerical_shard_tensor_check(mesh, module, [sharded_x], {}, check_grads=True) sous le marqueur multigpu_static, lancé comme

torchrun --nproc-per-node 4 -m pytest test/... --multigpu-static -m multigpu_static

Un test forward uniquement ne prouve presque rien — le gradient de poids est où vivent les bugs de sharding (il est Partial sur la maille de domaine et doit être réduit). Toujours check_grads=True, toujours désactiver TF32 pour la comparaison.

Ressources connexes

  • references/integration-checklist.md — checklist étape par étape pour rétrofitter un script d'entraînement/inférence existant, plus la matrice smoke 4-GPU valant la peine d'être scriptée.
  • references/new-op-patterns.md — anatomie des patches, niveaux d'enregistrement, et quel patch existant copier pour chaque classe d'opération.
  • references/testing.md — bootstrapping de test multi-GPU, numerical_shard_tensor_check, marqueurs, et invocation torchrun.
  • physicsnemo-discover — pour choisir les modèles, datapipes et exemples.

Skills similaires