You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在PyTorch迁移学习中快速简便实现Stochastic Depth(随机深度)

Quick Stochastic Depth Implementation for PyTorch Transfer Learning

Great question! In PyTorch, you absolutely can implement Stochastic Depth in a way that's just as quick and straightforward as using nn.Dropout—no need to rewrite entire ResNet blocks from scratch. Here's how to integrate it seamlessly into your transfer learning pipeline:

If you're using torchvision 0.13 or newer, there's an official stochastic_depth utility you can wrap into a reusable module (mirroring the simplicity of nn.Dropout):

import torch
import torch.nn as nn
from torchvision.ops import stochastic_depth

class StochasticDepth(nn.Module):
    def __init__(self, p: float = 0.5, mode: str = "row"):
        super().__init__()
        self.p = p
        # "row" = skip blocks per sample; "batch" = skip entire batch of blocks
        self.mode = mode

    def forward(self, x):
        # Automatically disables random skipping during evaluation
        return stochastic_depth(x, self.p, self.mode, self.training)

2. Add It to Your Transfer Learning Model

Once you have the module, you can insert it into your pre-trained model just like you would with nn.Dropout. For ResNets (the most common use case for Stochastic Depth), attach it to the residual branch of each block:

from torchvision.models import resnet50

# Load pre-trained ResNet50
model = resnet50(pretrained=True)

# Inject Stochastic Depth into every residual block's residual branch
for name, module in model.named_modules():
    # Target the sequential layers forming the residual branch (skip downsample layers)
    if isinstance(module, nn.Sequential) and "layer" in name and "downsample" not in name:
        module.add_module("stochastic_depth", StochasticDepth(p=0.2, mode="row"))

# Adjust the classifier head for your task (e.g., 10-class classification)
model.fc = nn.Linear(2048, 10)

Key Practical Notes

  • Train/Eval Behavior: Just like nn.Dropout, the module automatically turns off random skipping during evaluation, returning the full block output (scaled to match training expectations).
  • Hyperparameters: Use p (probability of dropping a residual block) between 0.1-0.5 for most tasks. The "row" mode is more standard, as it adds per-sample regularization.
  • Transfer Learning Best Practices: If freezing pre-trained layers, only add Stochastic Depth to custom classifier layers or new residual blocks. If fine-tuning the entire model, apply it to all blocks for stronger regularization.

Custom Implementation (For Older TorchVision Versions)

If you're on an older torchvision version, here's a minimal, bug-free custom module that behaves identically to the official one:

import torch
import torch.nn as nn

class StochasticDepth(nn.Module):
    def __init__(self, p: float = 0.5, mode: str = "row"):
        super().__init__()
        self.p = p
        self.mode = mode
        assert mode in ["row", "batch"], "Mode must be either 'row' or 'batch'"

    def forward(self, x):
        if not self.training:
            return x
        
        survival_prob = 1.0 - self.p
        # Create mask shape based on selected mode
        shape = (x.shape[0],) + (1,) * (x.ndim - 1) if self.mode == "row" else (1,) * x.ndim
        
        # Generate Bernoulli mask and scale output to preserve training expectation
        mask = torch.empty(shape, dtype=x.dtype, device=x.device).bernoulli_(survival_prob)
        return x * mask / survival_prob

This module works exactly like nn.Dropout—initialize it once, drop it into your model wherever you need it, and you're ready to train!

内容的提问来源于stack exchange,提问作者CuCaRot

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.29 17:42:36