zachlopez/sample_1
0
1#! /usr/bin/env python32# coding=utf-83 4# This code is licensed under a non-commercial license.5 6import os7import sys8import argparse9from tqdm import trange10 11import torch12import torch.nn.functional as F13import numpy as np14from IPython import embed15from operator import add16from style_utils import to_var, top_k_logits17import pickle18import csv19 20from gpt2tunediscrim import ClassificationHead21 22#lab_root = os.path.join(os.path.abspath(os.path.dirname(__file__)), '..', '..')23#sys.path.insert(1, lab_root)24 25from pytorch_pretrained_bert import GPT2LMHeadModel, GPT2Tokenizer26 27SmallConst = 1e-1528enc = GPT2Tokenizer.from_pretrained('gpt-2_pt_models/345M/')29 30def perturb_past(past, model, prev, args, classifier, good_index=None, stepsize=0.01, vocab_size=50257,31 original_probs=None, accumulated_hidden=None, true_past=None, grad_norms=None):32 window_length = args.window_length33 gm_scale, kl_scale = args.fusion_gm_scale, args.fusion_kl_scale34 one_hot_vectors = []35 for good_list in good_index:36 good_list = list(filter(lambda x: len(x) <= 1, good_list))37 good_list = torch.tensor(good_list).cuda()38 num_good = good_list.shape[0]39 one_hot_good = torch.zeros(num_good, vocab_size).cuda()40 one_hot_good.scatter_(1, good_list, 1)41 one_hot_vectors.append(one_hot_good)42 43 44 # Generate inital perturbed past45 past_perturb_orig = [(np.random.uniform(0.0, 0.0, p.shape).astype('float32'))46 for p in past]47 48 if accumulated_hidden is None:49 accumulated_hidden = 050 51 if args.decay:52 decay_mask = torch.arange(0., 1.0 + SmallConst, 1.0/(window_length))[1:]53 else:54 decay_mask = 1.055 56 # Generate a mask is gradient perturbated is based on a past window57 _, _, _, current_length, _ = past[0].shape58 59 if current_length > window_length and window_length > 0:60 ones_key_val_shape = tuple(past[0].shape[:-2]) + tuple([window_length]) + tuple(61 past[0].shape[-1:])62 63 zeros_key_val_shape = tuple(past[0].shape[:-2]) + tuple([current_length - window_length]) + tuple(64 past[0].shape[-1:])65 66 ones_mask = torch.ones(ones_key_val_shape)67 ones_mask = decay_mask*ones_mask.permute(0, 1, 2, 4, 3)68 ones_mask = ones_mask.permute(0, 1, 2, 4, 3)69 70 window_mask = torch.cat((ones_mask, torch.zeros(zeros_key_val_shape)), dim=-2).cuda()71 else:72 window_mask = torch.ones_like(past[0]).cuda()73 74 loss_per_iter = []75 for i in range(args.num_iterations):76 past_perturb = [torch.from_numpy(p_) for p_ in past_perturb_orig]77 past_perturb = [to_var(p_, requires_grad=True) for p_ in past_perturb]78 79 perturbed_past = list(map(add, past, past_perturb))80 81 _, _, _, current_length, _ = past_perturb[0].shape82 83 # Compute hidden using perturbed past84 _, future_past = model(prev, past=perturbed_past)85 hidden = model.hidden_states86 new_accumulated_hidden = accumulated_hidden + torch.sum(hidden, dim=1).detach()87 88 # TODO: Check the layer-norm consistency of this with trained discriminator89 logits = model.forward_hidden(hidden)90 logits = logits[:, -1, :]91 probabs = F.softmax(logits, dim=-1)92 loss = 0.093 loss_list = []94 if args.loss_type == 1 or args.loss_type == 3:95 for one_hot_good in one_hot_vectors:96 good_logits = torch.mm(probabs, torch.t(one_hot_good))97 loss_word = good_logits98 loss_word = torch.sum(loss_word)99 loss_word = -torch.log(loss_word)100 #loss_word = torch.sum(loss_word) /torch.sum(one_hot_good)101 loss += loss_word102 loss_list.append(loss_word)103 print('words', loss.data.cpu().numpy())104 105 if args.loss_type == 2 or args.loss_type == 3:106 ce_loss = torch.nn.CrossEntropyLoss()107 new_true_past = true_past108 for i in range(args.horizon_length):109 110 future_probabs = F.softmax(logits, dim=-1) # Get softmax111 future_probabs = torch.unsqueeze(future_probabs, dim=1)112 113 _, new_true_past = model(future_probabs, past=new_true_past)114 future_hidden = model.hidden_states # Get expected hidden states115 new_accumulated_hidden = new_accumulated_hidden + torch.sum(future_hidden, dim=1)116 117 predicted_sentiment = classifier(new_accumulated_hidden / (current_length + 1 + args.horizon_length))118 119 label = torch.tensor([args.label_class], device='cuda', dtype=torch.long)120 discrim_loss = ce_loss(predicted_sentiment, label)121 print('discrim', discrim_loss.data.cpu().numpy())122 loss += discrim_loss123 loss_list.append(discrim_loss)124 125 126 kl_loss = 0.0127 if kl_scale > 0.0:128 p = (F.softmax(original_probs[:, -1, :], dim=-1))129 p = p + SmallConst * (p <= SmallConst).type(torch.FloatTensor).cuda().detach()130 correction = SmallConst * (probabs <= SmallConst).type(torch.FloatTensor).cuda().detach()131 corrected_probabs = probabs + correction.detach()132 kl_loss = kl_scale * ((corrected_probabs * (corrected_probabs / p).log()).sum())133 #print('kl_loss', kl_loss.data.cpu().numpy())134 loss += kl_loss # + discrim_loss135 136 print((loss - kl_loss).data.cpu().numpy())137 138 loss_per_iter.append(loss.data.cpu().numpy())139 loss.backward()140 if grad_norms is not None and args.loss_type == 1:141 grad_norms = [torch.max(grad_norms[index], torch.norm(p_.grad * window_mask)) for index, p_ in142 enumerate(past_perturb)]143 else:144 grad_norms = [(torch.norm(p_.grad * window_mask) + SmallConst) for index, p_ in enumerate(past_perturb)]145 146 grad = [147 -stepsize * (p_.grad * window_mask / grad_norms[index] ** args.gamma).data.cpu().numpy()148 for index, p_ in enumerate(past_perturb)]149 past_perturb_orig = list(map(add, grad, past_perturb_orig))150 151 for p_ in past_perturb:152 p_.grad.data.zero_()153 154 new_past = []155 for p in past:156 new_past.append(p.detach())157 158 past = new_past159 160 past_perturb = [torch.from_numpy(p_) for p_ in past_perturb_orig]161 past_perturb = [to_var(p_, requires_grad=True) for p_ in past_perturb]162 perturbed_past = list(map(add, past, past_perturb))163 164 return perturbed_past, new_accumulated_hidden, grad_norms, loss_per_iter165 166 167def latent_perturb(model, args, context=None, sample=True, device='cuda'):168 if args.discrim == 'clickbait':169 classifier = ClassificationHead(class_size=2, embed_size=1024).to(device)170 classifier.load_state_dict(torch.load("discrim_models/clickbait_classifierhead.pt"))171 classifier.eval()172 args.label_class = 1 # clickbaity173 174 elif args.discrim == 'sentiment':175 classifier = ClassificationHead(class_size=5, embed_size=1024).to(device)176 classifier.load_state_dict(torch.load("discrim_models/sentiment_classifierhead.pt"))177 classifier.eval()178 if args.label_class < 0:179 raise Exception('Wrong class for sentiment, use --label-class 2 for *very positive*, 3 for *very negative*')180 #args.label_class = 2 # very pos181 #args.label_class = 3 # very neg182 183 elif args.discrim == 'toxicity':184 classifier = ClassificationHead(class_size=2, embed_size=1024).to(device)185 classifier.load_state_dict(torch.load("discrim_models/toxicity_classifierhead.pt"))186 classifier.eval()187 args.label_class = 0 # not toxic188 else:189 classifier = None190 191 # Get tokens for the list of positive words192 def list_tokens(word_list):193 token_list = []194 for word in word_list:195 token_list.append(enc.encode(" " + word))196 return token_list197 198 199 good_index = []200 if args.bag_of_words:201 bags_of_words = args.bag_of_words.split(";")202 for wordlist in bags_of_words:203 with open(wordlist, "r") as f:204 words = f.read()205 words = words.split('\n')206 good_index.append(list_tokens(words))207 208 if args.bag_of_words and classifier:209 print('Both PPLM-BoW and PPLM-Discrim are on. This is not optimized.')210 args.loss_type = 3211 212 elif args.bag_of_words:213 args.loss_type = 1214 print('Using PPLM-BoW')215 216 elif classifier is not None:217 args.loss_type = 2218 print('Using PPLM-Discrim')219 220 else:221 raise Exception('Supply either --bag-of-words (-B) or --discrim -D')222 223 224 original, _, _ = sample_from_hidden(model=model, args=args, context=context, device=device,225 perturb=False, good_index=good_index, classifier=classifier)226 torch.cuda.empty_cache()227 228 perturbed_list = []229 discrim_loss_list = []230 loss_in_time_list = []231 232 for i in range(args.num_samples):233 perturbed, discrim_loss, loss_in_time = sample_from_hidden(model=model, args=args, context=context,234 device=device, perturb=True, good_index=good_index,235 classifier=classifier)236 perturbed_list.append(perturbed)237 if classifier is not None:238 discrim_loss_list.append(discrim_loss.data.cpu().numpy())239 loss_in_time_list.append(loss_in_time)240 241 torch.cuda.empty_cache()242 243 244 return original, perturbed_list, discrim_loss_list, loss_in_time_list245 246 247def sample_from_hidden(model, args, classifier, context=None, past=None, device='cuda',248 sample=True, perturb=True, good_index=None):249 output = torch.tensor(context, device=device, dtype=torch.long).unsqueeze(0) if context else None250 251 grad_norms = None252 loss_in_time = []253 for i in trange(args.length, ascii=True):254 255 # Get past/probs for current output, except for last word256 # Note that GPT takes 2 inputs: past + current-token257 # Therefore, use everything from before current i/p token to generate relevant past258 259 260 if past is None and output is not None:261 prev = output[:, -1:]262 _, past = model(output[:, :-1])263 original_probs, true_past = model(output)264 true_hidden = model.hidden_states265 266 else:267 original_probs, true_past = model(output)268 true_hidden = model.hidden_states269 270 # Modify the past if necessary271 272 if i >= args.grad_length:273 current_stepsize = args.stepsize * 0274 else:275 current_stepsize = args.stepsize276 277 if not perturb or args.num_iterations == 0:278 perturbed_past = past279 280 else:281 accumulated_hidden = model.hidden_states[:, :-1, :]282 accumulated_hidden = torch.sum(accumulated_hidden, dim=1)283 284 perturbed_past, _, grad_norms, loss_per_iter = perturb_past(past, model, prev, args,285 good_index=good_index, stepsize=current_stepsize,286 original_probs=original_probs,287 true_past=true_past,288 accumulated_hidden=accumulated_hidden,289 classifier=classifier,290 grad_norms=grad_norms)291 loss_in_time.append(loss_per_iter)292 293 test_logits, past = model(prev, past=perturbed_past)294 # test_logits = F.softmax(test_logits[:, -1, :], dim=-1)295 # likelywords = torch.topk(test_logits, k=10, dim=-1)296 # print(enc.decode(likelywords[1].tolist()[0]))297 298 if classifier is not None:299 ce_loss = torch.nn.CrossEntropyLoss()300 predicted_sentiment = classifier(torch.mean(true_hidden, dim=1))301 label = torch.tensor([args.label_class], device='cuda', dtype=torch.long)302 true_discrim_loss = ce_loss(predicted_sentiment, label)303 print("true discrim loss", true_discrim_loss.data.cpu().numpy())304 else:305 true_discrim_loss = 0 306 307 hidden = model.hidden_states # update hidden308 logits = model.forward_hidden(hidden)309 logits = logits[:, -1, :] / args.temperature # + SmallConst310 311 # logits = top_k_logits(logits, k=args.top_k) # + SmallConst312 313 log_probs = F.softmax(logits, dim=-1)314 315 # Fuse the modified model and original model316 if perturb:317 318 # original_probs = top_k_logits(original_probs[:, -1, :]) #+ SmallConst319 original_probs = F.softmax(original_probs[:, -1, :], dim=-1)320 # likelywords = torch.topk(original_probs, k=10, dim=-1)321 # print(enc.decode(likelywords[1].tolist()[0]))322 323 gm_scale = args.fusion_gm_scale324 log_probs = ((log_probs ** gm_scale) * (original_probs ** (1 - gm_scale))) # + SmallConst325 326 log_probs = top_k_logits(log_probs, k=args.top_k, probs=True) # + SmallConst327 328 if torch.sum(log_probs) <= 1:329 log_probs = log_probs / torch.sum(log_probs)330 331 else:332 logits = top_k_logits(logits, k=args.top_k) # + SmallConst333 log_probs = F.softmax(logits, dim=-1)334 335 if sample:336 # likelywords = torch.topk(log_probs, k=args.top_k, dim=-1)337 # print(enc.decode(likelywords[1].tolist()[0]))338 # print(likelywords[0].tolist())339 prev = torch.multinomial(log_probs, num_samples=1)340 else:341 _, prev = torch.topk(log_probs, k=1, dim=-1)342 # if perturb:343 # prev = future344 output = prev if output is None else torch.cat((output, prev), dim=1) # update output345 print(enc.decode(output.tolist()[0]))346 347 return output, true_discrim_loss, loss_in_time348 349 350def run_model():351 parser = argparse.ArgumentParser()352 parser.add_argument('--model_path', '-M', type=str, default='gpt-2_pt_models/345M/',353 help='pretrained model name or path to local checkpoint')354 parser.add_argument('--bag-of-words', '-B', type=str, default=None, 355 help='Bags of words used for PPLM-BoW. Multiple BoWs separated by ;')356 parser.add_argument('--discrim', '-D', type=str, default=None, 357 choices=('clickbait', 'sentiment', 'toxicity'), 358 help='Discriminator to use for loss-type 2')359 parser.add_argument('--label-class', type=int, default=-1, help='Class label used for the discriminator')360 parser.add_argument('--stepsize', type=float, default=0.02)361 parser.add_argument("--length", type=int, default=100)362 parser.add_argument("--seed", type=int, default=0)363 parser.add_argument("--temperature", type=float, default=1.0)364 parser.add_argument("--top_k", type=int, default=10)365 parser.add_argument("--fusion-gm-scale", type=float, default=0.9)366 parser.add_argument("--fusion-kl-scale", type=float, default=0.01)367 parser.add_argument('--nocuda', action='store_true', help='no cuda')368 parser.add_argument('--uncond', action='store_true', help='Generate from end-of-text as prefix')369 parser.add_argument("--cond-text", type=str, default='The lake', help='Prefix texts to condition on')370 parser.add_argument('--num-iterations', type=int, default=3)371 parser.add_argument('--grad-length', type=int, default=10000)372 parser.add_argument('--num-samples', type=int, default=1,373 help='Number of samples to generate from the modified latents')374 parser.add_argument('--horizon-length', type=int, default=1, help='Length of future to optimize over')375 # parser.add_argument('--force-token', action='store_true', help='no cuda')376 parser.add_argument('--window-length', type=int, default=0,377 help='Length of past which is being optimizer; 0 corresponds to infinite window length')378 parser.add_argument('--decay', action='store_true', help='whether to decay or not')379 parser.add_argument('--gamma', type=float, default=1.5)380 381 args = parser.parse_args()382 383 torch.manual_seed(args.seed)384 np.random.seed(args.seed)385 386 device = 'cpu' if args.nocuda else 'cuda'387 388 model = GPT2LMHeadModel.from_pretrained(args.model_path)389 model.to(device)390 model.eval()391 392 # Freeze GPT-2 weights393 for param in model.parameters():394 param.requires_grad = False395 pass396 397 if args.uncond:398 seq = [[50256, 50256]]399 400 else:401 raw_text = args.cond_text402 while not raw_text:403 print('Did you forget to add `--cond-text`? ')404 raw_text = input("Model prompt >>> ")405 seq = [[50256] + enc.encode(raw_text)]406 407 collect_gen = dict()408 current_index = 0 409 for out in seq:410 411 text = enc.decode(out)412 print("=" * 40 + " Prefix of sentence " + "=" * 40)413 print(text)414 print("=" * 80)415 416 out1, out_perturb, discrim_loss_list, loss_in_time_list = latent_perturb(model=model, args=args, context=out,417 device=device)418 419 text_whole = enc.decode(out1.tolist()[0])420 421 print("=" * 80)422 print("=" * 40 + " Whole sentence (Original)" + "=" * 40)423 print(text_whole)424 print("=" * 80)425 426 out_perturb_copy = out_perturb427 428 generated = 0429 for out_perturb in out_perturb_copy:430 try:431 print("=" * 40 + " Whole sentence (Perturbed)" + "=" * 40)432 text_whole = enc.decode(out_perturb.tolist()[0])433 print(text_whole)434 print("=" * 80)435 except:436 pass437 collect_gen[current_index] = [out, out_perturb, out1]438 # Save the prefix, perturbed seq, original seq for each index439 440 current_index = current_index + 1441 442 return443 444 445if __name__ == '__main__':446 run_model()447 