pytorch-fsdp2 skill
Adds PyTorch FSDP2 (fully_shard) to training scripts with correct init, sharding, mixed precision/offload config, and distributed checkpointing. Use when models exceed single-GPU memory or when you need DTensor-based sharding with DeviceMesh.
Is the pytorch-fsdp2 skill safe?
Clean: nothing in its files matched our rules. We read 13 files in the folder on 2026-09-28.
No findings.
Install the pytorch-fsdp2 skill
A skill is a folder. Copy it into your agent's skills folder and the agent loads it when the task matches its description.
git clone --depth 1 https://github.com/Orchestra-Research/AI-Research-SKILLs.git /tmp/AI-Research-SKILLs mkdir -p ~/.claude/skills cp -r /tmp/AI-Research-SKILLs/08-distributed-training/pytorch-fsdp2 ~/.claude/skills/pytorch-fsdp2
In the Claude apps, zip the folder and upload it from the Skills settings. The folder on GitHub
The instructions your agent would load
SKILL.md as published, without the frontmatter. Read it on GitHub
Skill: Use PyTorch FSDP2 (fully_shard) correctly in a training script
This skill teaches a coding agent how to add PyTorch FSDP2 to a training loop with correct initialization, sharding, mixed precision/offload configuration, and checkpointing.
FSDP2 in PyTorch is exposed primarily via torch.distributed.fsdp.fullyshard and the FSDPModule methods it adds in-place to modules. See: references/pytorchfullyshardapi.md, references/pytorchfsdp2tutorial.md.
When to use this skill
Use FSDP2 when:
- Your model doesn’t fit on one GPU (parameters + gradients + optimizer state).
- You want an eager-mode sharding approach that is DTensor-based per-parameter sharding (more inspectable, simpler sharded state dicts) than FSDP1.
- You may later compose DP with Tensor Parallel using DeviceMesh.
Avoid (or be careful) if:
- You need strict backwards-compatible checkpoints across PyTorch versions (DCP warns against this).
- You’re forced onto older PyTorch versions without the FSDP2 stack.
Alternatives (when FSDP2 is not the best fit)
- DistributedDataParallel (DDP): Use the standard data-parallel wrapper when you want classic distributed data parallel training.
- FullyShardedDataParallel (FSDP1): Use the original FSDP wrapper for parameter sharding across data-parallel workers.
Reference: references/pytorchddpnotes.md, references/pytorchfsdp1api.md.
Contract the agent must follow
- Launch with torchrun and set the CUDA device per process (usually via LOCAL_RANK).
- Apply fullyshard() bottom-up**, i.e., shard submodules (e.g., Transformer blocks) before the root module.
- Call model(input), not model.forward(input), so the FSDP2 hooks run (unless you explicitly unshard() or register the forward method).
- Create the optimizer after sharding and make sure it is built on the DTensor parameters (post-fully_shard).
- Checkpoint using Distributed Checkpoint (DCP) or the distributed-state-dict helpers, not naïve torch.save(model.state_dict()) unless you deliberately gather to full tensors.
(Each of these rules is directly described in the official API docs/tutorial; see references.)
Step-by-step procedure
0) Version & environment sanity
- Prefer a recent stable PyTorch where the docs show FSDP2 and DCP updated recently.
- Use torchrun --nprocpernode ... and ensure RANK, WORLDSIZE, LOCALRANK are visible.
Reference: references/pytorchfsdp2tutorial.md (launch commands and setup), references/pytorchfullyshard_api.md (user contract).
1) Initialize distributed and set device
Minimal, correct pattern:
- dist.initprocessgroup(backend="nccl")
- torch.cuda.setdevice(int(os.environ["LOCALRANK"]))
- Optionally create a DeviceMesh to describe the data-parallel group(s)
Reference: references/pytorchdevicemesh_tutorial.md (why DeviceMesh exists & how it manages process groups).
2) Build model on meta device (recommended for very large models)
For big models, initialize on meta, apply sharding, then materialize weights on GPU:
- with torch.device("meta"): model = ...
- apply fullyshard(...) on submodules, then fullyshard(model)
- model.to_empty(device="cuda")
- model.reset_parameters() (or your init routine)
Reference: references/pytorchfsdp2tutorial.md (migration guide shows this flow explicitly).
3) Apply fully_shard() bottom-up (wrapping policy = “apply where needed”)
Do not only call fully_shard on the topmost module.
Recommended sharding pattern for transformer-like models:
- iterate modules, if isinstance(m, TransformerBlock): fully_shard(m, ...)
- then fully_shard(model, ...)
Why:
- fully_shard forms “parameter groups” for collective efficiency and excludes params already grouped by earlier calls. Bottom-up gives better overlap and lower peak memory.
Reference: references/pytorchfullyshard_api.md (bottom-up requirement and why).
4) Configure reshardafterforward for memory/perf trade-offs
Default behavior:
- None means True for non-root modules and False for root modules (good default).
Heuristics:
- If you’re memory-bound: keep defaults or force True on many blocks.
- If you’re throughput-bound and can afford memory: consider keeping unsharded params longer (root often False).
- Advanced: use an int to reshard to a smaller mesh after forward (e.g., intra-node) if it’s a meaningful divisor.
Reference: references/pytorchfullyshard_api.md (full semantics).
5) Mixed precision & offload (optional but common)
FSDP2 uses:
- mppolicy=MixedPrecisionPolicy(paramdtype=..., reducedtype=..., outputdtype=..., castforwardinputs=...)
- offload_policy=CPUOffloadPolicy() if you want CPU offload
Rules of thumb:
- Start with BF16 parameters/reductions on H100/A100-class GPUs (if numerically stable for your model).
- Keep reduce_dtype aligned with your gradient reduction expectations.
- If you use CPU offload, budget for PCIe/NVLink traffic and runtime overhead.
Reference: references/pytorchfullyshard_api.md (MixedPrecisionPolicy / OffloadPolicy classes).
6) Optimizer, gradient clipping, accumulation
- Create the optimizer after sharding so it holds DTensor params.
- If you need gradient accumulation / no_sync:
- use the FSDP2 mechanism (setrequiresgradientsync) instead of FSDP1’s nosync().
Gradient clipping:
- Use the approach shown in the FSDP2 tutorial (“Gradient Clipping and Optimizer with DTensor”), because parameters/gradients are DTensors.
Reference: references/pytorchfsdp2tutorial.md.
7) Checkpointing: prefer DCP or distributed state dict helpers
Two recommended approaches:
A) Distributed Checkpoint (DCP) — best default
- DCP saves/loads from multiple ranks in parallel and supports load-time resharding.
- DCP produces multiple files (often at least one per rank) and operates “in place”.
B) Distributed state dict helpers
- getmodelstatedict / setmodelstatedict with StateDictOptions(fullstatedict=True, cpuoffload=True, broadcastfrom_rank0=True, ...)
- For optimizer: getoptimizerstatedict / setoptimizerstatedict
Avoid:
- Saving DTensor state dicts with plain torch.save unless you intentionally convert with DTensor.full_tensor() and manage memory carefully.
References:
- references/pytorchdcpoverview.md (DCP behavior and caveats)
- references/pytorchdcprecipe.md and references/pytorchdcpasync_recipe.md (end-to-end usage)
- references/pytorchfsdp2tutorial.md (DTensor vs DCP state-dict flows)
- references/pytorchexamplesfsdp2.md (working checkpoint scripts)
More skills from Orchestra-Research/AI-Research-SKILLs
- Aacademic-plottingGenerates publication-quality figures for ML papers from research context. Given a paper section or description, extracts system components and relationships to generate architecture diagrams via Gemini. Given experiment results or data, auto-selects chart type and generates data-driven figures via matplotlib/seaborn. Use when creating any figure for a conference paper.
- Aara-compilerCompiles any research input — PDF papers, GitHub repositories, experiment logs, code directories, or raw notes — into a complete Agent-Native Research Artifact (ARA) with cognitive layer (claims, concepts, heuristics), physical layer (configs, code stubs), exploration graph, and grounded evidence. Use when ingesting a paper or codebase into a structured, machine-executable knowledge package, building an ARA from scratch, or converting research outputs into a falsifiable, agent-traversable form.
- Aara-research-managerRecords research provenance as a post-task epilogue, scanning conversation history at the end of a coding or research session to extract decisions, experiments, dead ends, claims, heuristics, and pivots, and writing them into the ara/ directory with user-vs-AI provenance tags. Use as a session epilogue — never during execution — to maintain a faithful, auditable trace of how a research project actually evolved.
- Aara-rigor-reviewerPerforms ARA Seal Level 2 semantic epistemic review on Agent-Native Research Artifacts, scoring six dimensions (evidence relevance, falsifiability, scope calibration, argument coherence, exploration integrity, methodological rigor) and producing a constructive, severity-ranked report with a Strong Accept-to-Reject recommendation. Use after Level 1 structural validation passes, when an ARA needs an objective epistemic critique before publication or release.
- Aaudiocraft-audio-generationPyTorch library for audio generation including text-to-music (MusicGen) and text-to-sound (AudioGen). Use when you need to generate music from text descriptions, create sound effects, or perform melody-conditioned music generation.
- Aautogpt-agentsAutonomous AI agent platform for building and deploying continuous agents. Use when creating visual workflow agents, deploying persistent autonomous agents, or building complex multi-step AI automation systems.
- AautoresearchOrchestrates end-to-end autonomous AI research projects using a two-loop architecture. The inner loop runs rapid experiment iterations with clear optimization targets. The outer loop synthesizes results, identifies patterns, and steers research direction. Routes to domain-specific skills for execution, supports continuous agent operation via Claude Code /loop and OpenClaw heartbeat, and produces research presentations and papers. Use when starting a research project, running autonomous experiments, or managing a multi-hypothesis research effort.
- Aawq-quantizationActivation-aware weight quantization for 4-bit LLM compression with 3x speedup and minimal accuracy loss. Use when deploying large models (7B-70B) on limited GPU memory, when you need faster inference than GPTQ with better accuracy preservation, or for instruction-tuned and multimodal models. MLSys 2024 Best Paper Award winner.
- CaxolotlExpert guidance for fine-tuning LLMs with Axolotl - YAML configs, 100+ models, LoRA/QLoRA, DPO/KTO/ORPO/GRPO, multimodal support
- Ablip-2-vision-languageVision-language pre-training framework bridging frozen image encoders and LLMs. Use when you need image captioning, visual question answering, image-text retrieval, or multimodal chat with state-of-the-art zero-shot performance.
- Abrainstorming-research-ideasGuides researchers through structured ideation frameworks to discover high-impact research directions. Use when exploring new problem spaces, pivoting between projects, or seeking novel angles on existing work.
- AchromaOpen-source embedding database for AI applications. Store embeddings and metadata, perform vector and full-text search, filter by metadata. Simple 4-function API. Scales from notebooks to production clusters. Use for semantic search, RAG applications, or document retrieval. Best for local development and open-source projects.