Skip to content
This repository was archived by the owner on Jul 13, 2026. It is now read-only.

Commit 23976ea

Browse files
authored
Add LitModular.gradient_modifier (#266)
Copy gradient modifier usage from Adversary to LitModular for training universal perturbations.
1 parent 8e4ac26 commit 23976ea

1 file changed

Lines changed: 12 additions & 0 deletions

File tree

‎mart/models/modular.py‎

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,7 @@ def __init__(
2525
self,
2626
modules,
2727
optimizer,
28+
gradient_modifier=None,
2829
lr_scheduler=None,
2930
training_sequence=None,
3031
training_step_log=None,
@@ -71,6 +72,7 @@ def __init__(
7172
# Set bias_decay and norm_decay to 0.
7273
self.optimizer_fn = OptimizerFactory(self.optimizer_fn)
7374

75+
self.gradient_modifier = gradient_modifier
7476
self.lr_scheduler = lr_scheduler
7577

7678
# Be backwards compatible by turning list into dict where each item is its own key-value
@@ -112,6 +114,16 @@ def configure_optimizers(self):
112114

113115
return configure_optimizers(self.model, self.optimizer_fn, self.lr_scheduler)
114116

117+
def configure_gradient_clipping(
118+
self, optimizer, gradient_clip_val=None, gradient_clip_algorithm=None
119+
):
120+
# Configuring gradient clipping in pl.Trainer is still useful, so use it.
121+
super().configure_gradient_clipping(optimizer, gradient_clip_val, gradient_clip_algorithm)
122+
123+
if self.gradient_modifier:
124+
for group in optimizer.param_groups:
125+
self.gradient_modifier(group["params"])
126+
115127
def forward(self, **kwargs):
116128
return self.model(**kwargs)
117129

0 commit comments

Comments
 (0)