PyTorch中冻结中间层参数的深度解析与实践

PyTorch中冻结中间层参数的深度解析与实践

本教程深入探讨了在PyTorch中冻结神经网络特定中间层参数的两种常见方法:torch.no_grad()上下文管理器和设置参数的requires_grad = False属性。文章通过代码示例详细阐述了两种方法的原理、效果及适用场景,并明确指出requires_grad = False是实现精确中间层冻结的推荐方案,同时提供了验证层是否被冻结的技巧,旨在帮助开发者准确控制模型训练过程中的参数更新。

在深度学习模型训练过程中,我们经常会遇到需要冻结模型中某些层(即不更新这些层的参数)而只训练其他层的场景,例如在迁移学习中冻结预训练模型的特征提取层,或者在多任务学习中只更新特定任务相关的层。本文将详细探讨pytorch中实现这一目标的方法。

理解参数冻结的原理

在PyTorch中,参数更新是通过反向传播计算梯度并由优化器应用到参数上的。冻结一个层意味着阻止其参数参与梯度计算和随后的更新。这通常通过控制参数的requires_grad属性来实现。当requires_grad为False时,PyTorch的自动求导引擎将不会为该参数计算梯度,从而阻止其被优化器更新。

方法一:使用 torch.no_grad() 上下文管理器

torch.no_grad()是一个上下文管理器,它会禁用在其作用域内所有操作的梯度计算。这意味着,任何在with torch.no_grad():块中执行的操作,都不会构建计算图,也不会跟踪梯度。

让我们通过一个简单的三层线性网络为例来演示:

import torchimport torch.nn as nnimport torch.optim as optim# 定义一个简单的模型class SimpleModel(nn.Module):    def __init__(self):        super(SimpleModel, self).__init__()        self.lin0 = nn.Linear(1, 2)        self.lin1 = nn.Linear(2, 2)        self.lin2 = nn.Linear(2, 10)    def forward_with_no_grad(self, x):        x = self.lin0(x)        with torch.no_grad():            x = self.lin1(x) # 尝试冻结lin1        x = self.lin2(x)        return x# 实例化模型model_no_grad = SimpleModel()# 记录初始参数initial_lin0_weight = model_no_grad.lin0.weight.clone()initial_lin1_weight = model_no_grad.lin1.weight.clone()initial_lin2_weight = model_no_grad.lin2.weight.clone()# 模拟训练步骤input_data = torch.randn(1, 1)target = torch.randint(0, 10, (1,))criterion = nn.CrossEntropyLoss()optimizer = optim.SGD(model_no_grad.parameters(), lr=0.01)print("--- 使用 torch.no_grad() 冻结中间层 ---")print("初始 lin0 权重:n", initial_lin0_weight)print("初始 lin1 权重:n", initial_lin1_weight)print("初始 lin2 权重:n", initial_lin2_weight)# 前向传播与反向传播output = model_no_grad.forward_with_no_grad(input_data)loss = criterion(output, target)optimizer.zero_grad()loss.backward()optimizer.step()# 检查参数变化print("n训练后 lin0 权重:n", model_no_grad.lin0.weight)print("训练后 lin1 权重:n", model_no_grad.lin1.weight)print("训练后 lin2 权重:n", model_no_grad.lin2.weight)# 验证是否冻结print("nlin0 权重是否变化:", not torch.equal(initial_lin0_weight, model_no_grad.lin0.weight))print("lin1 权重是否变化:", not torch.equal(initial_lin1_weight, model_no_grad.lin1.weight))print("lin2 权重是否变化:", not torch.equal(initial_lin2_weight, model_no_grad.lin2.weight))

分析 torch.no_grad() 的效果:上述代码运行后会发现,lin0和lin1的参数都没有更新,而只有lin2的参数发生了变化。这是因为当lin1的操作在torch.no_grad()块中执行时,其输出张量x(来自lin1)的grad_fn属性将为None,这意味着从lin1往前的计算图被截断了。因此,尽管lin2的梯度可以正常计算并回传到lin1的输出,但由于lin1的操作没有梯度跟踪,导致无法计算lin1自身的梯度,也无法将梯度继续回传到lin0。最终结果是,lin0和lin1的参数都不会得到更新。

结论: torch.no_grad() 适用于冻结整个模型或模型的一部分,使其在推理阶段不消耗内存来存储梯度信息,或者在训练时完全禁用某些部分的梯度更新。但它不适合精确地冻结中间层而允许其上游层更新的场景。

方法二:设置 requires_grad = False

这是在PyTorch中实现精确层冻结的推荐方法。通过将特定层的参数的requires_grad属性设置为False,我们可以明确告诉PyTorch的自动求导引擎不需要为这些参数计算梯度。

import torchimport torch.nn as nnimport torch.optim as optim# 定义一个简单的模型class SimpleModel(nn.Module):    def __init__(self):        super(SimpleModel, self).__init__()        self.lin0 = nn.Linear(1, 2)        self.lin1 = nn.Linear(2, 2)        self.lin2 = nn.Linear(2, 10)    def forward(self, x):        x = self.lin0(x)        x = self.lin1(x)        x = self.lin2(x)        return x# 实例化模型model_requires_grad = SimpleModel()# 冻结lin1层的参数model_requires_grad.lin1.weight.requires_grad = Falsemodel_requires_grad.lin1.bias.requires_grad = False# 记录初始参数initial_lin0_weight_rg = model_requires_grad.lin0.weight.clone()initial_lin1_weight_rg = model_requires_grad.lin1.weight.clone()initial_lin2_weight_rg = model_requires_grad.lin2.weight.clone()# 注意:优化器只应传入 requires_grad 为 True 的参数optimizer_rg = optim.SGD(filter(lambda p: p.requires_grad, model_requires_grad.parameters()), lr=0.01)# 模拟训练步骤input_data = torch.randn(1, 1)target = torch.randint(0, 10, (1,))criterion = nn.CrossEntropyLoss()print("n--- 使用 requires_grad = False 冻结中间层 ---")print("初始 lin0 权重:n", initial_lin0_weight_rg)print("初始 lin1 权重:n", initial_lin1_weight_rg)print("初始 lin2 权重:n", initial_lin2_weight_rg)# 前向传播与反向传播output = model_requires_grad(input_data)loss = criterion(output, target)optimizer_rg.zero_grad()loss.backward()optimizer_rg.step()# 检查参数变化print("n训练后 lin0 权重:n", model_requires_grad.lin0.weight)print("训练后 lin1 权重:n", model_requires_grad.lin1.weight)print("训练后 lin2 权重:n", model_requires_grad.lin2.weight)# 验证是否冻结print("nlin0 权重是否变化:", not torch.equal(initial_lin0_weight_rg, model_requires_grad.lin0.weight))print("lin1 权重是否变化:", not torch.equal(initial_lin1_weight_rg, model_requires_grad.lin1.weight))print("lin2 权重是否变化:", not torch.equal(initial_lin2_weight_rg, model_requires_grad.lin2.weight))

分析 requires_grad = False 的效果:运行上述代码后,你会发现lin0和lin2的参数都得到了更新,而只有lin1的参数保持不变。这是因为:

lin1.weight.requires_grad = False和lin1.bias.requires_grad = False明确地告诉PyTorch不要为这些参数计算梯度。在反向传播时,尽管梯度会流经lin1,但由于lin1的参数被标记为不需要梯度,PyTorch会跳过其梯度计算,并继续将梯度回传到lin0。优化器在初始化时,通过filter(lambda p: p.requires_grad, model_requires_grad.parameters())确保它只接收那些requires_grad=True的参数进行更新。

结论: requires_grad = False 是实现精确冻结模型中特定层(包括中间层)的正确且推荐的方法。它允许梯度流经被冻结的层,但不会更新该层自身的参数,同时能将梯度正确地传递给更上游的层。

验证层是否被冻结

在实际操作中,可以通过以下几种方式来验证层是否成功被冻结:

检查 param.requires_grad 属性:在设置后,可以打印出model.lin1.weight.requires_grad来确认其是否为False。

检查 param.grad 属性:在执行loss.backward()之后,检查被冻结层的参数(例如model.lin1.weight.grad)是否为None。如果为None,则表示没有为该参数计算梯度。

检查参数值是否变化:在训练循环开始前记录参数的初始值,经过一个或多个训练步骤后,再次检查这些参数的值。如果参数值未发生变化,则说明该层已被冻结。这正是本文示例代码中采用的方法。

总结与最佳实践

精确冻结中间层: 始终使用设置参数的requires_grad = False属性来冻结模型中的特定层。优化器初始化: 当冻结部分层时,务必在初始化优化器时,只将那些requires_grad = True的参数传递给优化器。例如:optimizer = torch.optim.SGD(filter(lambda p: p.requires_grad, model.parameters()), lr=0.01)。torch.no_grad() 的适用场景: torch.no_grad() 主要用于推理阶段,或者在训练过程中完全禁用某一部分的梯度计算,它会截断计算图,不适合需要梯度回传到上游层的场景。模型状态: 冻结层与model.train()和model.eval()没有直接冲突。model.eval()主要影响nn.BatchNorm和nn.Dropout等层在训练和评估模式下的行为,而requires_grad控制的是参数是否更新。

通过理解和正确应用requires_grad = False,开发者可以灵活地控制PyTorch模型中各层的训练状态,从而实现更复杂的训练策略,例如微调预训练模型或进行部分模型的更新。

以上就是PyTorch中冻结中间层参数的深度解析与实践的详细内容,更多请关注创想鸟其它相关文章!

版权声明:本文内容由互联网用户自发贡献,该文观点仅代表作者本人。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。
如发现本站有涉嫌抄袭侵权/违法违规的内容, 请发送邮件至 chuangxiangniao@163.com 举报,一经查实,本站将立刻删除。
发布者:程序猿,转转请注明出处:https://www.chuangxiangniao.com/p/1368386.html

赞 (0)
打赏 微信扫一扫 微信扫一扫 支付宝扫一扫 支付宝扫一扫
PyTorch中精确冻结中间层参数的策略与实践
上一篇 2025年12月14日 08:44:21
Python API请求指南:正确获取与解析API响应
下一篇 2025年12月14日 08:44:30

相关推荐

  • Krita中如何导出AI生成的分层图片?保存多层图像的步骤

    .kra格式是保存AI分层图像的最佳选择,因其完整保留Krita特有的图层、蒙版、滤镜等编辑信息,确保后续修改不受限;若需跨软件协作,则应导出为PSD格式,尽管可能损失部分Krita专属功能,但兼容性最广;TIFF适合高质量印刷场景,但分层支持不稳定;OpenEXR适用于含深度、法线等通道的专业合成…

    2026年9月22日
    100
  • mysql怎么执行sql命令 mysql输入代码创建表详细步骤

    mysql怎么执行sql命令 mysql输入代码创建表详细步骤mysql怎么执行sql命令 mysql输入代码创建表详细步骤mysql怎么执行sql命令 mysql输入代码创建表详细步骤mysql怎么执行sql命令 mysql输入代码创建表详细步骤

    在mysql中执行sql并创建表的步骤如下:1.通过命令行或图形工具连接数据库,使用mysql -u 用户名 -p并输入密码登录;2.选择或创建数据库,用use database_name或create database语句;3.使用create table定义表结构,如字段名、数据类型、约束等,例…

    2026年9月22日 • 用户投稿
    100
  • laravel如何使用Pipeline模式处理复杂逻辑_Laravel Pipeline模式处理复杂逻辑方法

    Laravel Pipeline通过将复杂流程拆分为多个独立处理步骤,实现代码解耦与职责分离。以用户注册为例,可依次执行发送欢迎邮件、分配角色、记录日志等操作,每个步骤由单独类实现__invoke方法,通过Pipeline::send($user)->through([…])-&g…

    2026年9月22日
    200
  • Swift 3到5.1新特性整理

    tocSwift 5.1Swift 5.0Result类型Raw string自定义字符串插值动态可调用类型处理未来的枚举值从try?抹平嵌套可选检查整数是否为偶数字典compactMapValues()方法撤回的功能: 带条件的计数Swift 4.2CaseIterable协议警告和错误指令动态查…

    2026年9月22日
    000
  • AffinityDesigner如何导出AI生成的矢量图片?保存图像的步骤

    答案是选择合适的矢量格式并调整导出设置。在Affinity Designer中导出AI生成的矢量图时,应根据用途选择SVG(适用于Web)、PDF(适用于打印和跨平台分享)或EPS(适用于老旧系统);导出前需检查文本是否转曲、颜色模式是否正确,并优化路径与位图设置以平衡质量与文件大小;从其他AI工具…

    2026年9月22日
    000
  • php-gd怎么制作缩略图_php-gd生成高质量缩略图

    使用PHP-GD生成高质量缩略图需保持宽高比、选用imagecopyresampled进行重采样,并合理设置JPEG质量(80-95),同时处理PNG透明通道,避免图像失真或背景变黑。 使用 PHP-GD 制作高质量缩略图,核心在于正确处理图像缩放、保持宽高比、避免失真,并选择合适的图像质量参数。下…

    2026年9月22日
    000
  • MySQL安装后初始密码在哪里查看?

    MySQL安装后初始密码在哪里查看?MySQL安装后初始密码在哪里查看?MySQL安装后初始密码在哪里查看?MySQL安装后初始密码在哪里查看?

    mysql安装后的初始密码取决于安装方式和操作系统,通常可在错误日志中找到。1. 查看mysql错误日志:linux系统使用grep命令查找/var/log/mysqld.log或类似路径;windows系统在data目录下的hostname.err中搜索“temporary password”。2…

    2026年9月22日 • 用户投稿
    100
  • PHP日志记录怎么做_PHP中Monolog库实现灵活强大的日志系统

    Monolog是PHP中基于PSR-3标准的主流日志库,通过Composer安装后可轻松实现日志记录。使用Logger类创建实例并添加Handler(如StreamHandler写入文件、NativeMailerHandler邮件报警)来管理不同级别(debug、info、error等)日志输出,支…

    2026年9月22日
    200
  • 如何在RunwayML导出AI生成的4K图片?保存高清图像的教程

    要从RunwayML获得4K图像,需结合高分辨率生成设置与AI放大工具。首先在RunwayML中选择最高可用分辨率(如1024×1024或更高),并通过精细提示词和负面提示词优化生成质量;随后利用内置增强功能或外部AI放大工具(如Topaz Gigapixel AI、Upscayl)将图像…

    2026年9月22日
    100
  • mysql如何添加主键索引 mysql创建主键索引的步骤详解

    mysql如何添加主键索引 mysql创建主键索引的步骤详解mysql如何添加主键索引 mysql创建主键索引的步骤详解mysql如何添加主键索引 mysql创建主键索引的步骤详解mysql如何添加主键索引 mysql创建主键索引的步骤详解

    mysql中添加主键索引主要有三种方式:1. 创建新表时直接添加主键,可在列定义后使用primary key或在所有列定义后单独声明;2. 在已有表上通过alter table添加主键,需确保目标列非空且唯一,必要时先清洗数据;3. 添加复合主键,适用于多列组合才能唯一标识记录的情况。主键索引在in…

    2026年9月22日 • 用户投稿
    000
  • PHP数组如何定义和使用_PHP数组定义与使用详细教程

    PHP数组是存储和管理多个值的核心工具,支持索引、关联、混合及多维结构;通过方括号定义,可灵活访问、修改、添加或删除元素,并利用foreach高效遍历。 PHP数组是存储一系列值的强大工具,无论这些值是简单的数据项,还是更复杂的结构。它的核心思想就是把一堆相关的数据“打包”在一起,通过一个统一的名字…

    2026年9月22日
    000
  • 逛京东先人一步下单E人E本EBOOK X14 Air笔记本 补贴立省10%

    逛京东先人一步下单E人E本EBOOK X14 Air笔记本 补贴立省10%逛京东先人一步下单E人E本EBOOK X14 Air笔记本 补贴立省10%逛京东先人一步下单E人E本EBOOK X14 Air笔记本 补贴立省10%逛京东先人一步下单E人E本EBOOK X14 Air笔记本 补贴立省10%

    10月13日10:00,京东抢先首发e人e本全新力作——ebook x14 air ai轻薄笔记本电脑,以仅898克的极致轻盈机身和卓越的本地ai算力,重新定义高效移动办公新标准。新品官方定价7999元,京东首发期间可享国家补贴直降10%,实付仅需7199元,晒单再赠50元京东e卡,下单即送高品质内…

    2026年9月22日 • 用户投稿
    000
  • Pictory如何快速生成AI视频?从文本到AI视频的完整教程

    Pictory通过智能算法将文字脚本转化为专业AI视频,核心在于自动分析文本、匹配视觉素材、生成语音并初步剪辑。用户登录后选择“Script to Video”,粘贴结构清晰的脚本,AI会自动分割场景并推荐素材,支持手动调整场景划分、替换素材、上传自定义图片视频以增强品牌一致性。平台提供多语言AI语…

    2026年9月22日
    000
  • Flyway多数据库与CI/CD测试集成策略

    本文深入探讨了在CI/CD流程中,如何高效地配置Flyway以管理多数据库环境下的迁移,尤其关注集成测试场景。我们将比较使用真实数据库服务、Testcontainers以及Flyway自身多数据库配置的优劣,并提供关于分离生产与测试环境迁移脚本的实用策略,旨在确保开发、测试与生产环境的数据一致性与流…

    2026年9月22日
    100
  • Spring Boot自定义Kafka配置与动态Bean注册最佳实践

    本文探讨了在Spring Boot应用中通过自定义注解简化Kafka配置的挑战与解决方案。重点介绍了如何利用META-INF/spring.factories实现早期自动配置,并详细阐述了使用ImportBeanDefinitionRegistrar在应用上下文初始化早期动态注册Kafka生产者工厂…

    2026年9月22日
    100
  • mysql安装后怎么维护 mysql日常维护操作大全

    mysql安装后怎么维护 mysql日常维护操作大全mysql安装后怎么维护 mysql日常维护操作大全mysql安装后怎么维护 mysql日常维护操作大全mysql安装后怎么维护 mysql日常维护操作大全

    开启并分析慢查询日志以优化 sql 性能;2. 定期使用逻辑或物理方式备份数据并异地存储;3. 监控连接数和服务器资源,防止资源耗尽;4. 定期执行 analyze、optimize 和 check 表操作以维护表健康;5. 合理管理日志配置与清理策略。mysql 安装后的日常维护主要包括慢查询监控…

    2026年9月22日 • 用户投稿
    100
  • 深度解析蝴蝶号如何实现AI实景24小时无人直播

    深度解析蝴蝶号如何实现AI实景24小时无人直播深度解析蝴蝶号如何实现AI实景24小时无人直播深度解析蝴蝶号如何实现AI实景24小时无人直播深度解析蝴蝶号如何实现AI实景24小时无人直播

    蝴蝶号能实现ai实景24小时无人直播,主要靠智能中控系统+实景画面采集+自动化互动机制。一、ai中控系统作为“大脑”,自动控制画面切换、语音播报、商品推荐和评论区互动,具备一定判断能力,确保稳定性与持续性。二、实景画面采集作为“眼睛”,通过高清摄像头和云台控制,在门店、仓库等场景采集实时画面,保障真…

    2026年9月22日 • 用户投稿
    200
  • 在Java中如何开发简易问答社区

    答案是Java结合Spring Boot可快速构建问答社区,通过设计questions、answers、users三张表实现数据存储,使用JPA进行持久化,前端用HTML+JS调用后端API完成用户提问、回答、查看与互动功能。 开发一个简易问答社区,核心是实现用户提问、回答、查看问题和互动功能。Ja…

    2026年9月22日
    100
  • HitPawVideoEditor如何制作AI视频?教你快速创建AI内容的步骤

    答案是HitPaw Video Editor通过AI文本转视频、AI图片生成、智能抠图、自动字幕等功能,显著提升视频创作效率。它以“AI创作+人工精修”模式降低制作门槛,帮助用户快速生成初稿、丰富视觉素材、简化复杂操作,并支持快速迭代,但需避免过度依赖AI,仍需人工打磨以确保情感表达与叙事质量。 ☞…

    2026年9月22日
    000
  • linux系统下codeblocks控制台打印中文乱码[通俗易懂]

    linux系统下codeblocks控制台打印中文乱码[通俗易懂]linux系统下codeblocks控制台打印中文乱码[通俗易懂]linux系统下codeblocks控制台打印中文乱码[通俗易懂]linux系统下codeblocks控制台打印中文乱码[通俗易懂]

    大家好,很高兴再次和大家见面,我是你们的朋友全栈君。 在Linux系统下使用CodeBlocks时,如果在控制台中打印中文可能会遇到乱码问题。以下是解决这一问题的详细步骤: 首先,我们来看一下在Linux系统下安装CodeBlocks后,运行以下代码时出现的问题: #include #include…

    2026年9月22日 • 用户投稿
    600

发表回复

登录后才能评论
关注微信