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éthodesforward(), 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.py—wrap_ddp,shard_spatial_params_(sélecteur basé sur le nom pour pos_embed/RoPE),wrap_fsdp_spatialexamples/weather/stormcast/utils/parallel.py—ParallelHelperde productionexamples/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. Avecdomain_size > 1, compilez régionalement : patch-embed / norms et MLPs par bloc / tête, laissant l'attention eager. Avecdomain_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 lesShardTensorSpecs 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_metasordinaires). 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_functionsdansphysicsnemo/domain_parallel/shard_tensor.py; la couverture de régression vit danstest/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)
TypeError: unsupported operand type(s) for +: 'ShardTensor' and 'ShardTensor'n'est presque jamais l'erreur réelle. Les dunder binaires convertissent uneNotImplementedErrorinterne enNotImplemented, et CPython émet ce message générique, avalant la vraie trace. Remplacez temporairementx + ypartorch.add(x, y)pour surfacer l'exception véritable.- 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é. Utilisezscatter_tensor(..., requires_grad=True)ou threading les gradients par les paramètres. torch.autograd.gradfonctionne directement sur ShardTensors — c'est une fonction de passthrough autograd (s'exécute sur les vrais objets tenseurs sousDisableTorchFunctionSubclass). 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 monkeypatchingtorch.autograd.grad(par ex. pour logger les appels) casse le passthrough :handle_torch_functionpasse le module-globalgradrésolu au moment de l'appel, donc les lookups d'identité voient votre wrapper.- Seules certaines fonctions sont passthrough-safe (
register_hook,register_post_accumulate_grad_hook,retain_grad,torch.autograd.grad— voir_autograd_passthrough_functionsdansshard_tensor.py). Toute autre méthode sensible à l'identité peut agir sur un temporaire converti. - 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. 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 :
- 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) dansCommDebugMode. - 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__), etShardTensor.register_named_function_handler("lib.op.default", wrapper)pour lestorch.library.custom_ops. - 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(explicitautograd.Functionavec 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.