如何在Flux.jl中冻结指定层参数用于迁移学习?
Flux.jl 迁移学习指定冻结层的实现方式
方式1:直接指定可训练参数集合
这种是最通用的方式,适配所有Flux版本,核心逻辑是只把需要更新的层的参数传给优化器,未传入的参数不会被梯度更新,自然处于冻结状态。
示例代码:
using Flux, Metalhead # 加载预训练模型 model = Metalhead.ResNet18(pretrain=true) # 仅提取最后一层全连接层fc的参数作为可训练参数 trainable_params = Flux.params(model.fc) # 训练流程和普通训练一致,仅需将trainable_params传入训练函数 opt = Adam(1e-4) loss(x, y) = Flux.crossentropy(model(x), y) Flux.train!(loss, trainable_params, dataloader, opt)
方式2:通过trainable属性标记冻结层
Flux 0.13及以上版本支持给单个层设置trainable属性,值为false时该层参数会被自动排除在可训练参数集合外,适合需要冻结多层的场景,操作更简洁。
示例代码:
using Flux, Metalhead # 加载预训练模型 model = Metalhead.ResNet18(pretrain=true) # 先冻结所有层 for layer in model.layers layer.trainable = false end # 仅开放最后一层的可训练权限 model.fc.trainable = true # 直接传入全部模型参数即可,Flux会自动过滤trainable=false的层的参数 trainable_params = Flux.params(model) # 后续训练逻辑不变 opt = Adam(1e-4) loss(x, y) = Flux.crossentropy(model(x), y) Flux.train!(loss, trainable_params, dataloader, opt)
注意事项
- 如果你的任务和预训练模型的输出维度不一致,需要先替换最后一层为匹配你任务的输出层,再设置可训练属性
- 可以通过打印
length(Flux.params(model))验证可训练参数的数量是否符合预期,避免层名拼写错误导致冻结失效 - 嵌套结构的自定义模型也可以递归设置子模块的
trainable属性,逻辑和上述操作一致
内容的提问来源于stack exchange,提问作者logankilpatrick
相关产品推荐
相关产品推荐

