Opacus 致力于让 PyTorch 模型的私密训练在用户端只需最少量的代码更改。正如您可能通过阅读 README 和入门教程所了解的那样,Opacus 通过接收您的模型、数据加载器和优化器,并返回经过封装的对应项来完成此操作,这些封装后的对应项可以执行与隐私相关的功能。
虽然大多数常见模型都与 Opacus 兼容,但并非所有模型都兼容。
nn.ReLU、nn.Tanh 等)和冻结的模块(其参数的 requires_grad 设置为 False)都是兼容的。GradSampleModule 和 opacus.layers 提供的实现都具有此属性。BatchNorm)不适合 DP,因为样本的归一化值取决于其他样本,因此与 Opacus 不兼容。InstanceNorm)适合 DP,除了某些配置(例如,当 track_running_stats 为 On 时)。期望您记住所有这些并加以处理是不合理的。这就是 Opacus 提供 ModuleValidator 来处理此问题的原因。
ModuleValidator 内部¶ModuleValidator 类有两个主要的类方法 validate() 和 fix()。
顾名思义,validate() 通过确保给定模块处于训练模式并且是 GradSampleModule 类型(即,模块可以捕获每个样本的梯度)来验证其与 Opacus 的兼容性。更重要的是,此方法还会检查子模块及其配置是否存在兼容性问题(更多内容见下一节)。
fix() 方法尝试使模块与 Opacus 兼容。
在 Opacus 0.x 中,对每个支持模块的特定检查和必要的替换都通过一系列 if 检查在验证器中集中完成。添加新的验证检查和修复需要修改核心 Opacus 代码。在 Opacus 1.0 中,这已经模块化,允许您注册自己的自定义验证器和修复器。
在本教程的其余部分,我们将以 nn.BatchNorm 为例,并准确展示如何做到这一点。
我们知道 BatchNorm 模块不适合隐私保护,因此验证器应该抛出错误,例如这样
def validate_bathcnorm(module):
return [Exception("BatchNorm is not supported")]
要注册上述内容,您只需按如下方式装饰上述方法即可。
from opacus.validators import register_module_validator
@register_module_validator(
[nn.BatchNorm1d, nn.BatchNorm2d, nn.BatchNorm3d, nn.SyncBatchNorm]
)
def validate_bathcnorm(module):
return [Exception("BatchNorm is not supported")]
就是这样!上述操作将为所有这些模块注册 validate_bathcnorm():[nn.BatchNorm1d, nn.BatchNorm2d, nn.BatchNorm3d, nn.SyncBatchNorm],当您执行 privacy_engine.make_private() 时,此方法将与其他验证器一起自动调用。
该装饰器实质上是将您的方法添加到 ModuleValidator 的注册表中,以便在验证阶段循环使用。
只是一点小提示:建议您使您的验证异常尽可能清晰。Opacus 对上述情况的验证如下所示
@register_module_validator(
[nn.BatchNorm1d, nn.BatchNorm2d, nn.BatchNorm3d, nn.SyncBatchNorm]
)
def validate(module) -> None:
return [
ShouldReplaceModuleError(
"BatchNorm cannot support training with differential privacy. "
"The reason for it is that BatchNorm makes each sample's normalized value "
"depend on its peers in a batch, ie the same sample x will get normalized to "
"a different value depending on who else is in its batch. "
"Privacy-wise, this means that we would have to put a privacy mechanism there too. "
"While it can in principle be done, there are now multiple normalization layers that "
"do not have this issue: LayerNorm, InstanceNorm and their generalization GroupNorm "
"are all privacy-safe since they don't have this property."
"We offer utilities to automatically replace BatchNorms to GroupNorms and we will "
"release pretrained models to help transition, such as GN-ResNet ie a ResNet using "
"GroupNorm, pretrained on ImageNet"
)
]. # quite a mouthful, but is super clear! ;)
验证很好,但我们能在可能的情况下解决问题吗?答案当然是肯定的。语法与验证器的语法几乎相同。
例如,BatchNorm 可以用 GroupNorm 替换,而不会造成任何有意义的性能损失,并且仍然是隐私友好的。在 Opacus 中,我们这样做如下
def _batchnorm_to_groupnorm(module) -> nn.GroupNorm:
"""
Converts a BatchNorm ``module`` to GroupNorm module.
This is a helper function.
Args:
module: BatchNorm module to be replaced
Returns:
GroupNorm module that can replace the BatchNorm module provided
Notes:
A default value of 32 is chosen for the number of groups based on the
paper *Group Normalization* https://arxiv.org/abs/1803.08494
"""
return nn.GroupNorm(
min(32, module.num_features), module.num_features, affine=module.affine
)
from opacus.validators.utils import register_module_fixer
@register_module_fixer(
[nn.BatchNorm1d, nn.BatchNorm2d, nn.BatchNorm3d, nn.SyncBatchNorm]
)
def fix(module) -> nn.GroupNorm:
logger.info(
"The default batch_norm fixer replaces BatchNorm with GroupNorm."
" The batch_norm validator module also offers implementations to replace"
" it with InstanceNorm or Identity. Please check them out and override the"
" fixer if those are more suitable for your needs."
)
return _batchnorm_to_groupnorm(module)
当您调用 privacy_engine.make_private() 时,Opacus 不会自动为您修复模块;它期望模块在传入之前是兼容的。但是,这可以很容易地按如下方式完成
import torch
from opacus.validators import ModuleValidator
model = torch.nn.Linear(2,1)
if not ModuleValidator.is_valid(model):
model = ModuleValidator.fix(model)
如果您想使用自定义修复器来代替提供的修复器,您只需使用相同的装饰器来装饰您的函数即可。请注意,注册顺序很重要,最后注册的函数将被使用。
例如:要只用 InstanceNorm 替换 BatchNorm2d(同时使用 GroupNorm 作为 BatchNorm1d 和 BatchNorm3d 的默认替换),您可以这样做
import torch.nn as nn
from opacus.validators import register_module_fixer
@register_module_validator([nn.BatchNorm2d])
def fix_batchnorm2d(module):
return nn.InstanceNorm2d(module.num_features)
希望本教程对您有所帮助!欢迎您查看 opacus/validators/ 下的代码以获取详细信息。如果您有任何问题或意见,请随时在我们的 论坛 上发表。