Skip to content

Commit 3a171d4

Browse files
authored
Merge pull request #8 from matanbt/fix/soft-prompt-fp16-nan
Fix soft-prompt optimization diverging to NaN in half precision
2 parents 09d61f0 + 5d96be8 commit 3a171d4

3 files changed

Lines changed: 21 additions & 12 deletions

File tree

tropt/optimizer/gbda_optimizer.py

Lines changed: 8 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -146,17 +146,19 @@ def optimize_trigger(
146146
trigger_ids: Int[Tensor, "trigger_seq_len"] = tokenizer.encode_trigger(initial_trigger).to(self.model.device)
147147
vocab_size = self.model.vocab_size
148148
device = self.model.device
149+
model_dtype = self.model.dtype
149150
trigger_seq_len = trigger_ids.shape[0]
150151

151-
# Initialize logit matrix theta (here called `trigger_probs`)
152+
# Initialize logit matrix theta (here called `trigger_probs`).
153+
# stored in high precision for the optimizer
152154
if self.init_mode == "random":
153155
trigger_probs: Float[Tensor, "seq_len vocab_size"] = (
154-
torch.randn(trigger_seq_len, vocab_size, device=device, dtype=self.model.dtype)
156+
torch.randn(trigger_seq_len, vocab_size, device=device, dtype=torch.float32)
155157
* self.init_noise_scale
156158
)
157159
else: # "from_trigger"
158160
trigger_probs: Float[Tensor, "seq_len vocab_size"] = torch.zeros(
159-
trigger_seq_len, vocab_size, device=device, dtype=self.model.dtype,
161+
trigger_seq_len, vocab_size, device=device, dtype=torch.float32,
160162
)
161163
for i in range(trigger_seq_len):
162164
trigger_probs[i, trigger_ids[i]] = self.initial_coeff
@@ -181,7 +183,7 @@ def optimize_trigger(
181183
# (we repeat `trigger_probs_samples` so the gradient computation will draw multiple samples (w/ gumbel-softmax) from the same (optimized) distribution.)
182184
trigger_probs_samples = trigger_probs.unsqueeze(0).repeat(self.n_grad_samples, 1, 1)
183185
trigger_grad = self.model.compute_grad_from_tokens(
184-
candidate_trigger_probs=trigger_probs_samples, # (n_grad_samples, trigger_seq_len, vocab_size)
186+
candidate_trigger_probs=trigger_probs_samples.to(model_dtype), # (n_grad_samples, trigger_seq_len, vocab_size)
185187
loss_func=self.loss_func,
186188
do_gumbel_softmax=True,
187189
gumbel_softmax_temp=temperature,
@@ -191,7 +193,7 @@ def optimize_trigger(
191193
avg_grad = trigger_grad.mean(dim=0) # -> (trigger_seq_len, vocab_size)
192194

193195
# take grad step:
194-
trigger_probs.grad = avg_grad
196+
trigger_probs.grad = avg_grad.float()
195197
if self.grad_clip_norm is not None:
196198
torch.nn.utils.clip_grad_norm_([trigger_probs], self.grad_clip_norm)
197199
optimizer.step()
@@ -239,7 +241,7 @@ def optimize_trigger(
239241
best_trigger_ids=final_best_ids,
240242
losses=best.losses,
241243
trigger_strs=best.trigger_strs,
242-
best_trigger_probs=trigger_probs.detach(),
244+
best_trigger_probs=trigger_probs.detach().to(model_dtype),
243245
)
244246

245247
return result

tropt/optimizer/pez_optimizer.py

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -85,6 +85,8 @@ def optimize_trigger(
8585

8686
# Initialize continuous embeddings from the initial trigger tokens
8787
trigger_embeds = self.model._embedding_layer(trigger_ids.unsqueeze(0)) # (1, trigger_seq_len, embed_dim)
88+
model_dtype = trigger_embeds.dtype
89+
trigger_embeds = trigger_embeds.float() # stored in high precision for the optimizer
8890

8991
# Initialize optimizer on continuous embeddings
9092
optimizer = self.GDOptimizer(
@@ -98,7 +100,7 @@ def optimize_trigger(
98100

99101
# Forward projection: project continuous embeddings to nearest vocab tokens
100102
projected_ids, projected_embeds = self._project_to_vocab(
101-
trigger_embeds.squeeze(0), embedding_matrix
103+
trigger_embeds.squeeze(0).to(model_dtype), embedding_matrix
102104
)
103105

104106
# Compute gradient w.r.t. the projected embeddings.
@@ -113,7 +115,7 @@ def optimize_trigger(
113115
curr_loss = curr_loss.item()
114116

115117
# Set gradient on the continuous embeddings and step
116-
trigger_embeds.grad = trigger_grad
118+
trigger_embeds.grad = trigger_grad.float()
117119
optimizer.step()
118120

119121
# Decode current discrete trigger for tracking
@@ -125,7 +127,7 @@ def optimize_trigger(
125127

126128
# Final projection and evaluation on discrete tokens
127129
final_ids, _ = self._project_to_vocab(
128-
trigger_embeds.squeeze(0), embedding_matrix
130+
trigger_embeds.squeeze(0).to(model_dtype), embedding_matrix
129131
)
130132
final_trigger_str = tokenizer.decode_trigger(final_ids)
131133
final_loss = self.model.compute_loss_from_tokens(

tropt/optimizer/soft_optimizer.py

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -71,6 +71,8 @@ def optimize_trigger(
7171
trigger_ids = tokenizer.encode_trigger(initial_trigger).to(self.model.device)
7272

7373
trigger_embeds = self.model._embedding_layer(trigger_ids.unsqueeze(0)) # (1, trigger_seq_len, embd_dim)
74+
model_dtype = trigger_embeds.dtype
75+
trigger_embeds = trigger_embeds.float() # stored in high precision for the optimizer
7476

7577
# Initialize the optimizer on the trigger embeddings
7678
optimizer = self.GDOptimizer([trigger_embeds], lr=self.learning_rate)
@@ -83,19 +85,22 @@ def optimize_trigger(
8385
# Compute gradients w.r.t. trigger embeddings
8486
trigger_grad, curr_loss = self.model.compute_grad_from_embeds(
8587
loss_func=self.loss_func,
86-
candidate_trigger_embeds=trigger_embeds,
88+
candidate_trigger_embeds=trigger_embeds.to(model_dtype),
8789
normalize_grads=False,
8890
return_loss=True,
8991
) # grad: (1, trigger_seq_len, embed_dim); loss: (1,)
9092
curr_loss = curr_loss.item()
9193

9294
# Set gradient on trigger embeddings
93-
trigger_embeds.grad = trigger_grad
95+
trigger_embeds.grad = trigger_grad.float()
9496

9597
# Adam step
9698
optimizer.step()
9799

98-
best.update(loss=curr_loss, trigger_emb=trigger_embeds.detach().squeeze(0))
100+
best.update(
101+
loss=curr_loss,
102+
trigger_emb=trigger_embeds.detach().squeeze(0).to(model_dtype),
103+
)
99104
self.log(loss=curr_loss, lr=optimizer.param_groups[0]["lr"], grad_norm=trigger_grad.norm().item())
100105

101106
result = best.to_result()

0 commit comments

Comments
 (0)