@@ -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
0 commit comments