diff --git a/tropt/optimizer/gbda_optimizer.py b/tropt/optimizer/gbda_optimizer.py index 9ed064c..62a35f0 100644 --- a/tropt/optimizer/gbda_optimizer.py +++ b/tropt/optimizer/gbda_optimizer.py @@ -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 @@ -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, @@ -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() @@ -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 diff --git a/tropt/optimizer/pez_optimizer.py b/tropt/optimizer/pez_optimizer.py index 1ecfec9..a3b24cf 100644 --- a/tropt/optimizer/pez_optimizer.py +++ b/tropt/optimizer/pez_optimizer.py @@ -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( @@ -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. @@ -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 @@ -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( diff --git a/tropt/optimizer/soft_optimizer.py b/tropt/optimizer/soft_optimizer.py index 6872609..6bfe788 100644 --- a/tropt/optimizer/soft_optimizer.py +++ b/tropt/optimizer/soft_optimizer.py @@ -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) @@ -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()