Run Distributed Full-Model DPO with PyTorch FSDP

October 11, 2026 • guides

Direct Preference Optimization (DPO) aligns autoregressive language models to human preferences without requiring a separate reward model or the unstable dynamics of reinforcement learning with PPO. Instead of training an auxiliary critic to score outputs, DPO reformulates the policy objective so the language model's own sequence likelihoods are optimized directly against a static reference baseline. For teams shipping conversational agents, code assistants, or domain-tuned models, full-parameter DPO has become the standard post-training phase for reducing refusal behaviours, suppressing hallucinations, and reinforcing required stylistic constraints.

Running full-parameter alignment across multi-billion-parameter architectures places severe pressure on compute and memory. DPO evaluates both a trainable active policy and an immutable reference model over pairs of chosen and rejected responses in each step, so memory use is higher than ordinary supervised fine-tuning. Model weights, gradients, optimizer state, and activations all compete for the same accelerator memory before a single backward pass completes.

The implementation below is adapted from PyTorch torchtune's full_dpo_distributed.py, distributed under the BSD-3-Clause licence. It uses PyTorch Fully Sharded Data Parallel (FSDP), memory-efficient activation management, and concatenated tensor batching so you can run full DPO on accessible multi-GPU nodes without exhausting VRAM. The source file supports FSDP sharding, CPU offload, activation checkpointing, activation offloading, bf16 or fp32 precision, gradient accumulation, checkpoint resumption, and several logging backends.

Prerequisites

Before executing the distributed training recipe, ensure your environment matches the underlying library and infrastructure requirements.

  • Hardware: Single node containing 1 to 8 CUDA or XPU devices. If the GPU does not support bfloat16, the recipe falls back to fp32.
  • PyTorch: PyTorch 2.5 or later if you want non-blocking activation offloading streams; otherwise synchronous activation offloading is still available.
  • Dependencies: torchtune, torchdata, omegaconf, and tqdm.
  • CPU memory: When fsdp_cpu_offload is enabled, host RAM must be large enough to hold offloaded parameters, gradients, and optimizer states.
  • Checkpoints: An SFT base checkpoint and tokenizer, plus an identical or compatible reference model checkpoint. The reference checkpoint is loaded separately from the policy checkpoint.

Step 1: Sharding the Active Policy with FSDP and Activation Checkpointing

In distributed DPO, the active policy updates its parameters from loss gradients derived from preference log-probabilities. Keeping unpartitioned parameters on every GPU quickly causes out-of-memory failures. The recipe instantiates model definitions directly on the meta device to bypass CPU-to-GPU allocation bottlenecks, then shards parameters across ranks with PyTorch FSDP.

The _setup_model method handles instantiation, activation checkpointing through TransformerSelfAttentionLayer wrapping, optional activation offloading to host CPU memory, and parameter ingestion from state dictionaries without overloading rank memory.

    def _setup_model(
        self,
        cfg_model: DictConfig,
        enable_activation_checkpointing: bool,
        enable_activation_offloading: bool,
        fsdp_cpu_offload: bool,
        reshard_after_forward: bool,
        model_state_dict: dict[str, Any],
        custom_sharded_layers: Optional[list[str]] = None,
    ) -> nn.Module:
        """
        Model initialization has some important considerations:
           a. To minimize GPU peak memory, we initialize the model on meta device with
              the right dtype
           b. All ranks calls ``load_state_dict`` without peaking CPU RAMs since
              full state dicts are loaded with ``torch.load(mmap=True)``
        """

        utils.log_rank_zero(
            self._logger,
            "FSDP is enabled. Instantiating model and loading checkpoint on Rank 0 ...",
        )
        init_start = time.perf_counter()

        with training.set_default_dtype(self._dtype), torch.device("meta"):
            model = config.instantiate(cfg_model)

        if self._compile:
            training.compile_model(model, verbose=self._is_rank_zero)

        # original activation checkpointing (full) - flip the condition above
        if enable_activation_checkpointing:
            training.set_activation_checkpointing(
                model, auto_wrap_policy={modules.TransformerSelfAttentionLayer}
            )

        # For FSDP sharding
        fsdp_shard_conditions = [
            partial(
                training.get_shard_conditions,
                names_to_match=custom_sharded_layers,
            )
        ]
        training.shard_model(
            model=model,
            shard_conditions=fsdp_shard_conditions,
            cpu_offload=fsdp_cpu_offload,
            reshard_after_forward=reshard_after_forward,
        )

        with training.set_default_dtype(self._dtype), self._device:
            for m in model.modules():
                # RoPE is not covered in state dict
                if hasattr(m, "rope_init"):
                    m.rope_init()

        # This method will convert the full model state dict into a sharded state
        # dict and load into the model
        training.load_from_full_model_state_dict(
            model,
            model_state_dict,
            self._device,
            strict=True,
            cpu_offload=fsdp_cpu_offload,
        )

        # activation offloading
        self.activations_handling_ctx = training.get_act_offloading_ctx_manager(
            model, enable_activation_offloading
        )

        # Ensure no params and buffers are on meta device
        training.validate_no_params_on_meta_device(model)

        utils.log_rank_zero(
            self._logger,
            f"Instantiating model and loading checkpoint took {time.perf_counter() - init_start:.2f} secs",
        )

        if self._is_rank_zero:
            memory_stats = training.get_memory_stats(device=self._device)
            training.log_memory_stats(memory_stats)

        # disabling dropout if found - non-determinism leads to issues in e.g. comparing logprobs
        # between ref policy and current policy
        disable_dropout(model)

        # synchronize before training begins
        torch.distributed.barrier()

        return model

Initializing layers inside the torch.device("meta") context avoids allocating full physical memory buffers before sharding conditions run. Once shard_model decomposes the linear layers across distributed workers, load_from_full_model_state_dict populates only the local rank slice from disk. Calling disable_dropout(model) prevents stochastic mask variation from introducing differences between policy and reference forward evaluations.

Step 2: Instantiating and Freezing the Reference Baseline

DPO prevents policy collapse by regularizing the Kullback-Leibler divergence against the base SFT model. That requires evaluating every token sequence through an un-updated reference model. If configured incorrectly, maintaining two full models at once causes immediate GPU memory exhaustion.

_setup_reference_model creates this second policy, sets requires_grad = False on all tensors, and puts the module strictly into evaluation mode.

    def _setup_reference_model(
        self,
        cfg_model: DictConfig,
        fsdp_cpu_offload: bool,
        reshard_after_forward: bool,
        model_state_dict: dict[str, Any],
        custom_sharded_layers: Optional[list[str]] = None,
    ) -> nn.Module:
        """
        Similar to `self._setup_model`:
           a. To minimize GPU peak memory, we initialize the model on meta device with
              the right dtype
           b. All ranks calls ``load_state_dict`` without peaking CPU RAMs since
              full state dicts are loaded with ``torch.load(mmap=True)``

        Additionally, since the reference model is inference-only, we omit some training-specific
        optimizations.
        """

        utils.log_rank_zero(
            self._logger,
            "FSDP is enabled. Instantiating reference model and loading checkpoint on Rank 0 ...",
        )
        init_start = time.perf_counter()

        with training.set_default_dtype(self._dtype), torch.device("meta"):
            model = config.instantiate(cfg_model)

        if self._compile:
            training.compile_model(model, verbose=self._is_rank_zero)

        # For FSDP sharding
        fsdp_shard_conditions = [
            partial(
                training.get_shard_conditions,
                names_to_match=custom_sharded_layers,
            )
        ]
        training.shard_model(
            model=model,
            shard_conditions=fsdp_shard_conditions,
            cpu_offload=fsdp_cpu_offload,
            reshard_after_forward=reshard_after_forward,
        )

        with training.set_default_dtype(self._dtype), self._device:
            for m in model.modules():
                # RoPE is not covered in state dict
                if hasattr(m, "rope_init"):
                    m.rope_init()

        # This method will convert the full model state dict into a sharded state
        # dict and load into the model
        training.load_from_full_model_state_dict(
            model,
            model_state_dict,
            self._device,
            strict=True,
            cpu_offload=fsdp_cpu_offload,
        )

        # Ensure no params and buffers are on meta device
        training.validate_no_params_on_meta_device(model)

        utils.log_rank_zero(
            self._logger,
            f"Instantiating reference model and loading checkpoint took {time.perf_counter() - init_start:.2f} secs",
        )

        if self._is_rank_zero:
            memory_stats = training.get_memory_stats(device=self._device)
            training.log_memory_stats(memory_stats)

        # disabling dropout if found - non-determinism leads to issues in e.g. comparing logprobs
        # between ref policy and current policy
        disable_dropout(model)

        for p in model.parameters():
            p.requires_grad = False

        model.eval()

        # synchronize before training begins
        torch.distributed.barrier()

        return model

The reference model uses the same FSDP sharding conditions as the active model, but it omits activation checkpointing and offloading wrappers. Because backward differentiation is never called on the reference graph, intermediate activations do not need to be saved or recomputed.

Step 3: Preparing Preference Pairs with Stateful Distributed Sampling

DPO requires pairwise contrastive data containing a prompt and two distinct responses: a preferred completion y_w and a dispreferred completion y_l. The tokenizer formats both inputs into combined token sequences where prompt tokens are masked with CROSS_ENTROPY_IGNORE_IDX, typically -100.

The dataloader pipeline synchronizes sampling across ranks while remaining stateful, so interrupted runs can resume mid-epoch without reprocessing or skipping pairs.

    def _setup_data(
        self,
        cfg_dataset: DictConfig,
        shuffle: bool,
        batch_size: int,
        dataloader_state_dict: Optional[dict[str, Any]] = None,
    ) -> StatefulDataLoader:
        """
        All data related setup happens here. This recipe currently supports only
        map-style datasets. If a state_dict is provided (meaning we are resuming a training run),
        it is loaded into the dataloader.
        """

        if isinstance(cfg_dataset, ListConfig):
            datasets = [
                config.instantiate(single_cfg_dataset, tokenizer=self._tokenizer)
                for single_cfg_dataset in cfg_dataset
            ]
            ds = ConcatDataset(datasets=datasets)
        else:
            ds = config.instantiate(cfg_dataset, tokenizer=self._tokenizer)

        sampler = StatefulDistributedSampler(
            ds, num_replicas=self.world_size, rank=self.rank, shuffle=shuffle
        )

        dataloader = StatefulDataLoader(
            dataset=ds,
            batch_size=batch_size,
            sampler=sampler,
            # dropping last avoids shape issues with compile + flex attention
            drop_last=True,
            collate_fn=partial(
                padded_collate_dpo,
                padding_idx=self._tokenizer.pad_id,
                ignore_idx=CROSS_ENTROPY_IGNORE_IDX,
            ),
        )

        if dataloader_state_dict is not None:
            dataloader.load_state_dict(dataloader_state_dict)

        if self._is_rank_zero:
            self._logger.info("Dataset and Sampler are initialized.")

        return dataloader

padded_collate_dpo combines chosen and rejected prompt-response sequences into unified batches along the leading dimension. Setting drop_last=True enforces rigid tensor dimensions throughout the training loop, which prevents runtime graph recompilations when torch.compile is enabled.

Step 4: Computing Concatenated Forward Passes

Executing separate forward passes for chosen and rejected sequences doubles attention kernel dispatch overhead. Instead, modern DPO pipelines concatenate chosen and rejected examples into a single contiguous batch of shape [2B, L], dispatching one forward pass through the transformer blocks.

concatenated_forward runs this unified tensor, splits the output logits back into halves, and gathers sequence log-probabilities against masked target tokens.

    def concatenated_forward(
        self,
        model: nn.Module,
        batch: tuple[torch.Tensor, torch.Tensor],
        activations_handling: Optional[bool] = True,
    ) -> ChosenRejectedOutputs:
        """
        Run forward pass of the model with chosen and rejected samples concatenated.

        Args:
            model (nn.Module): The model to be used for the forward pass.
            batch (tuple[torch.Tensor, torch.Tensor]): tuple of input_ids and labels.

        Returns:
            Dataclass of chosen log probs, rejected log probs, chosen logits, rejected logits.
        """
        concatenated_input_ids, concatenated_labels = batch
        concatenated_input_ids = concatenated_input_ids.to(self._device)
        concatenated_labels = concatenated_labels.to(self._device)

        # formed by concatenating an equal number of "chosen" and "rejected".
        len_chosen = concatenated_input_ids.shape[0] // 2

        if activations_handling:
            with self.activations_handling_ctx:
                all_logits = model(concatenated_input_ids)
        else:
            all_logits = model(concatenated_input_ids)

        chosen_log_probs = rlhf.get_batch_log_probs(
            all_logits[:len_chosen],
            concatenated_labels[:len_chosen],
            return_average_logprobs=False,
        )

        rejected_log_probs = rlhf.get_batch_log_probs(
            all_logits[len_chosen:],
            concatenated_labels[len_chosen:],
            return_average_logprobs=False,
        )

        chosen_logits = all_logits[:len_chosen]
        rejected_logits = all_logits[len_chosen:]

        return ChosenRejectedOutputs(
            chosen_log_probs, rejected_log_probs, chosen_logits, rejected_logits
        )

The split len_chosen = concatenated_input_ids.shape[0] // 2 separates predictions for preferred and rejected completions. Passing un-averaged sequence log-probabilities into ChosenRejectedOutputs lets the loss compute length-normalized or sum-based implicit rewards.

Execution and Memory Management Strategies

Balancing throughput against GPU memory limits requires choosing the right sharding and offloading strategy. The table below summarizes the primary configuration paths supported by the torchtune DPO architecture.

Execution Mode Primary Configuration Memory Effect
Full FSDP reshard_after_forward: True Shards parameters, gradients, and optimizer state across ranks; resharding frees gathered weights after forward.
FSDP with CPU Offload fsdp_cpu_offload: True Moves parameters, gradients, and optimizer states to host RAM.
Activation Checkpointing enable_activation_checkpointing: True Drops forward activations and recomputes them during backward.
Activation Offloading enable_activation_offloading: True Moves activations to CPU and brings them back during backward; can overlap computation on PyTorch 2.5+.

Activation offloading requires activation checkpointing to be enabled. The recipe raises a runtime error if enable_activation_offloading is true while enable_activation_checkpointing is false. It also restricts activation offloading to CUDA or XPU devices.

Step 5: Initializing Distributed Environments and Entry Points

The launcher bootstraps the distributed backend and enforces multi-threading controls when offloading compute-intensive optimizer steps to CPU threads.

@config.parse
def recipe_main(cfg: DictConfig) -> None:
    """
    Entry point for the recipe.

    Configurable parameters are read in the following order:
        - Parameters specified in config (see available configs through ``tune ls``)
        - Overwritten by arguments from the command-line
    """
    if not training.is_distributed():
        raise RuntimeError(
            "Distributed finetune recipe should be run via a distributed launcher."
            "If using tune CLI, please specify --nnodes 1 and --nproc_per_node [num_gpus]"
        )
    if cfg.get("fsdp_cpu_offload", False):
        # Utilize all available CPU cores for intra-op parallelism. This provides ~2x
        # speed up when benchmarking fused AdamW on CPU
        training.set_torch_num_threads()

    config.log_config(recipe_name="FullDPORecipeDistributed", cfg=cfg)

    recipe = FullDPORecipeDistributed(cfg=cfg)
    recipe.setup(cfg=cfg)
    recipe.train()
    recipe.cleanup()


if __name__ == "__main__":
    sys.exit(recipe_main())

Invoke training through a multi-process launcher such as tune run --nproc_per_node 4 full_dpo_distributed --config llama3_2/8B_full_dpo. The launcher populates the environment variables consumed by init_process_group: RANK, WORLD_SIZE, MASTER_ADDR, and MASTER_PORT.

What to Watch Out For

Several subtle failure modes can undermine an alignment run.

First, verify that the reference model is truly frozen. The recipe sets requires_grad = False and calls model.eval(), but if you bypass _setup_reference_model or reuse an altered checkpoint, reference gradients will silently double memory use and corrupt the DPO objective.

Second, keep activation offloading constraints in mind. Offloading is only valid on CUDA or XPU, and it must be paired with activation checkpointing. Violating either condition raises a runtime error before training begins.

Third, never attempt mixed precision with fp16. The recipe raises ValueError: full fp16 training is not supported because DPO computes ratios of log-probabilities that can underflow or overflow outside bf16 or fp32 dynamic range. Use dtype=bf16 where supported, or dtype=fp32.

Fourth, treat raw logits as a memory spike source. The training loop computes policy_chosen_logits_mean and policy_rejected_logits_mean for metrics, then explicitly deletes the chosen and rejected logits before the reference forward pass. Preserve that pattern; removing the del statement causes peak memory to grow with vocabulary size.

Fifth, if you use optimizer-in-backward, do not combine it with gradient clipping or gradient accumulation. The recipe rejects such configurations because optimizer steps then occur inside the backward pass and those controls would conflict.

For a broader view of evaluation pitfalls that matter after alignment, see our Anthropic eval-gating report.

Where to Go Next

With the distributed DPO loop operational, consider these extensions. The training loop already checks loss_fn.is_reference_free; a reference-free objective skips the frozen reference model and eliminates that second sharded model entirely. Asynchronous checkpointing through enable_async_checkpointing can overlap large checkpoint writes with computation steps, reducing multi-node I/O stalls. If your deployment needs lower per-GPU memory, combine gradient_accumulation_steps with a smaller batch size before enabling activation checkpointing; the recipe makes total batch size equal to batch_size * number of GPUs * gradient_accumulation_steps.

Frequently asked questions

What is full DPO in torchtune?

Full DPO fine-tunes all parameters of a dense transformer against preference pairs, using a frozen reference model to regularize the policy. The FullDPORecipeDistributed recipe supports DPOLoss and RSOPLoss and avoids a separate reward model or PPO. It runs on 1 to 8 GPUs with FSDP sharding.

Does torchtune DPO support fp16?

No. Full fp16 training is explicitly disallowed by the recipe and raises a ValueError. Use bf16 or fp32 instead; if the GPU does not support bf16, the recipe falls back to fp32.

How does FSDP reduce memory during DPO?

FSDP shards parameters, gradients, and optimizer states across distributed workers instead of keeping full copies on each GPU. The recipe can also reshard parameters after the forward pass and optionally offload parameters, gradients, and optimizer states to CPU RAM.

How does concatenated_forward reduce DPO overhead?

It concatenates chosen and rejected sequences into a single batch of shape [2B, L], so the transformer runs one forward pass instead of two. The method then splits the logits back into chosen and rejected halves and returns their log-probabilities and logits for the loss.

Can I resume a torchtune DPO training run?

Yes. The recipe stores optimizer state, epochs run, global step, seed, and dataloader state, and it reloads them when resume_from_checkpoint is True. The stateful dataloader and StatefulDistributedSampler allow mid-epoch resume without reshuffling or skipping batches.

What losses does FullDPORecipeDistributed support?

The recipe supports DPOLoss and RSOPLoss, which is rejection sampling optimization. The training loop also checks loss_fn.is_reference_free, so a reference-free objective can skip the frozen reference model forward entirely.

Related Guides