ModuleValidator

class opacus.validators.module_validator.ModuleValidator[source]

Encapsulates all the validation logic required by Opacus. Also works as a namespace to hold registered validators and fixers.

classmethod fix(module, *, remove_inplace_ops=False, **kwargs)[source]

Make the module and sub_modules DP compatible by running registered custom fixers.

In addition to the per-module-type fixers, this always disables in-place activation flags (e.g. nn.ReLU(inplace=True)) since in-place writes are incompatible with Opacus’ backward hooks. This is a cheap, type-preserving fix that requires no tracing.

Functional/tensor in-place ops baked into forward (e.g. a ResNet’s out += identity) are only rewritten when remove_inplace_ops=True, because doing so requires symbolic tracing and changes the returned type. It is opt-in: models without such ops (or whose only in-place ops are the activation flags handled above) do not need it. If your model crashes under grad_sample_mode="hooks" with RuntimeError: Output 0 of BackwardHookFunction is a view and is being modified inplace, re-run with remove_inplace_ops=True.

When remove_inplace_ops=True and tracing succeeds, the returned object is a torch.fx.GraphModule, not the original nn.Module subclass. GraphModule is a subclass of nn.Module and preserves forward outputs for typical models, but isinstance checks against the original class, custom methods defined on the original class, pickling by class name, and similar type-dependent code will observe the GraphModule type instead. A warning is emitted in this case.

FX tracing also records a single execution path: data-independent control flow in the root forward (e.g. if self.training) is frozen to the branch taken during tracing, so a later .train()/.eval() toggle may have no effect. Standard nn.Dropout/nn.BatchNorm submodules keep their own flags and are unaffected.

Rewriting replaces in-place writes with out-of-place equivalents. Forward outputs are preserved, but aliasing semantics change: a tensor mutated in place is instead replaced by a new tensor, so other aliases to the original storage will not see the update. Code relying on in-place side effects via aliasing may diverge in behavior.

Models that cannot be symbolically traced (e.g. data-dependent control flow) are returned unchanged and the fallback is logged at INFO level.

Parameters:
  • module (Module) – The root module to be made compatible.

  • remove_inplace_ops (bool) – If True, rewrite functional/tensor in-place ops out-of-place via symbolic tracing. Defaults to False. No-op if the module cannot be traced.

  • **kwargs – Arbitrary keyword arguments forwarded to the per-module fixers.

Return type:

Module

Returns:

Fixed module. Type is torch.fx.GraphModule when tracing succeeds with remove_inplace_ops=True, otherwise the original (cloned) nn.Module type.

classmethod fix_and_validate(module, **kwargs)[source]

Fix the module and sub_modules first, and then run validation.

Parameters:
  • module (Module) – The root module to be fixed and validated

  • **kwargs – Arbitrary keyword arguments.

Return type:

Module

Returns:

Fixed module.

Raises:

UnsupportedModuleError in case of validation failures.

classmethod is_valid(module)[source]

Check if module and sub_modules are valid by running registered custom validators.

Parameters:

module (Module) – The root module to validate.

Return type:

bool

Returns:

bool

classmethod validate(module, *, strict=False)[source]

Validate module and sub_modules by running registered custom validators. Returns or raises exceptions depending on strict flag.

Parameters:
  • module (Module) – The root module to validate.

  • strict (bool) – Boolean to indicate whether to raise errors or return

  • errors. (the list of)

Raises:

UnsupportedModuleError in case of validation failures.

Return type:

List[UnsupportedModuleError]

opacus.validators.utils.register_module_fixer(target_class_or_classes, validator_class=<class 'opacus.validators.module_validator.ModuleValidator'>)[source]

Registers the decorated function as the fixer of target_class_or_classes, which is the function that will be invoked every time you want to fix an incompatoble module to make it work for training with Opacus. You may supply your own validator_class that holds the registry of FIXERS. The signature of every fixer is always the same:

>>> @register_module_fixer(MyCustomModel)
... def fix(module: nn.Module, **kwargs) -> nn.Module:
...    pass

It may help you to take a look at the existing fixers inside Opacus, under opacus.validators.

opacus.validators.utils.register_module_validator(target_class_or_classes, validator_class=<class 'opacus.validators.module_validator.ModuleValidator'>)[source]

Registers the decorated function as the validator of target_class_or_classes, which is the function that will be invoked every time you want to validate that a module is compatible for training with Opacus. You may supply your own validator_class that holds the registry of VALIDATORS. The signature of every validator is always the same:

>>> @register_module_validator(MyCustomModel)
... def validate(module: nn.Module, **kwargs) -> List[opacus.validators.errors.UnsupportedError]:
...    pass

It may help you to take a look at the existing validator inside Opacus, under opacus.validators.