Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 8 additions & 6 deletions tropt/optimizer/gbda_optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -146,17 +146,19 @@ def optimize_trigger(
trigger_ids: Int[Tensor, "trigger_seq_len"] = tokenizer.encode_trigger(initial_trigger).to(self.model.device)
vocab_size = self.model.vocab_size
device = self.model.device
model_dtype = self.model.dtype
trigger_seq_len = trigger_ids.shape[0]

# Initialize logit matrix theta (here called `trigger_probs`)
# Initialize logit matrix theta (here called `trigger_probs`).
# stored in high precision for the optimizer
if self.init_mode == "random":
trigger_probs: Float[Tensor, "seq_len vocab_size"] = (
torch.randn(trigger_seq_len, vocab_size, device=device, dtype=self.model.dtype)
torch.randn(trigger_seq_len, vocab_size, device=device, dtype=torch.float32)
* self.init_noise_scale
)
else: # "from_trigger"
trigger_probs: Float[Tensor, "seq_len vocab_size"] = torch.zeros(
trigger_seq_len, vocab_size, device=device, dtype=self.model.dtype,
trigger_seq_len, vocab_size, device=device, dtype=torch.float32,
)
for i in range(trigger_seq_len):
trigger_probs[i, trigger_ids[i]] = self.initial_coeff
Expand All @@ -181,7 +183,7 @@ def optimize_trigger(
# (we repeat `trigger_probs_samples` so the gradient computation will draw multiple samples (w/ gumbel-softmax) from the same (optimized) distribution.)
trigger_probs_samples = trigger_probs.unsqueeze(0).repeat(self.n_grad_samples, 1, 1)
trigger_grad = self.model.compute_grad_from_tokens(
candidate_trigger_probs=trigger_probs_samples, # (n_grad_samples, trigger_seq_len, vocab_size)
candidate_trigger_probs=trigger_probs_samples.to(model_dtype), # (n_grad_samples, trigger_seq_len, vocab_size)
loss_func=self.loss_func,
do_gumbel_softmax=True,
gumbel_softmax_temp=temperature,
Expand All @@ -191,7 +193,7 @@ def optimize_trigger(
avg_grad = trigger_grad.mean(dim=0) # -> (trigger_seq_len, vocab_size)

# take grad step:
trigger_probs.grad = avg_grad
trigger_probs.grad = avg_grad.float()
if self.grad_clip_norm is not None:
torch.nn.utils.clip_grad_norm_([trigger_probs], self.grad_clip_norm)
optimizer.step()
Expand Down Expand Up @@ -239,7 +241,7 @@ def optimize_trigger(
best_trigger_ids=final_best_ids,
losses=best.losses,
trigger_strs=best.trigger_strs,
best_trigger_probs=trigger_probs.detach(),
best_trigger_probs=trigger_probs.detach().to(model_dtype),
)

return result
8 changes: 5 additions & 3 deletions tropt/optimizer/pez_optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -82,6 +82,8 @@ def optimize_trigger(

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

# Initialize optimizer on continuous embeddings
optimizer = self.GDOptimizer(
Expand All @@ -95,7 +97,7 @@ def optimize_trigger(

# Forward projection: project continuous embeddings to nearest vocab tokens
projected_ids, projected_embeds = self._project_to_vocab(
trigger_embeds.squeeze(0), embedding_matrix
trigger_embeds.squeeze(0).to(model_dtype), embedding_matrix
)

# Compute gradient w.r.t. the projected embeddings.
Expand All @@ -110,7 +112,7 @@ def optimize_trigger(
curr_loss = curr_loss.item()

# Set gradient on the continuous embeddings and step
trigger_embeds.grad = trigger_grad
trigger_embeds.grad = trigger_grad.float()
optimizer.step()

# Decode current discrete trigger for tracking
Expand All @@ -122,7 +124,7 @@ def optimize_trigger(

# Final projection and evaluation on discrete tokens
final_ids, _ = self._project_to_vocab(
trigger_embeds.squeeze(0), embedding_matrix
trigger_embeds.squeeze(0).to(model_dtype), embedding_matrix
)
final_trigger_str = tokenizer.decode_trigger(final_ids)
final_loss = self.model.compute_loss_from_tokens(
Expand Down
11 changes: 8 additions & 3 deletions tropt/optimizer/soft_optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,8 @@ def optimize_trigger(
trigger_ids = tokenizer.encode_trigger(initial_trigger).to(self.model.device)

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

# Initialize the optimizer on the trigger embeddings
optimizer = self.GDOptimizer([trigger_embeds], lr=self.learning_rate)
Expand All @@ -83,19 +85,22 @@ def optimize_trigger(
# Compute gradients w.r.t. trigger embeddings
trigger_grad, curr_loss = self.model.compute_grad_from_embeds(
loss_func=self.loss_func,
candidate_trigger_embeds=trigger_embeds,
candidate_trigger_embeds=trigger_embeds.to(model_dtype),
normalize_grads=False,
return_loss=True,
) # grad: (1, trigger_seq_len, embed_dim); loss: (1,)
curr_loss = curr_loss.item()

# Set gradient on trigger embeddings
trigger_embeds.grad = trigger_grad
trigger_embeds.grad = trigger_grad.float()

# Adam step
optimizer.step()

best.update(loss=curr_loss, trigger_emb=trigger_embeds.detach().squeeze(0))
best.update(
loss=curr_loss,
trigger_emb=trigger_embeds.detach().squeeze(0).to(model_dtype),
)
self.log(loss=curr_loss, lr=optimizer.param_groups[0]["lr"], grad_norm=trigger_grad.norm().item())

result = best.to_result()
Expand Down
Loading