CoolFace
Apppublic

zachlopez/sample_1

sourceHugging Faceupdated 4y agoView on Hugging Face
0likes
pplm.py447 linesDownload Raw Back to paper_code
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