mirror of
https://github.com/wassname/Clover-Edition.git
synced 2026-09-09 11:13:26 +08:00
changed gpt-2 generator
This commit is contained in:
+2
-8
@@ -1,14 +1,8 @@
|
||||
import os
|
||||
from story.utils import *
|
||||
from google.cloud import storage
|
||||
import json
|
||||
from story.story_manager import *
|
||||
# from generator.web.web_generator import *
|
||||
# from generator.ctrl.ctrl_generator import *
|
||||
from generator.simple.simple_generator import *
|
||||
import tensorflow as tf
|
||||
import textwrap
|
||||
import sys
|
||||
from generator.gpt2.gpt2_generator import *
|
||||
|
||||
CRED_FILE = "./AI-Adventure-2bb65e3a4e2f.json"
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
*.pyc
|
||||
__pycache__
|
||||
@@ -0,0 +1,21 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2019 OpenAI
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
@@ -0,0 +1,28 @@
|
||||
import os
|
||||
import sys
|
||||
import requests
|
||||
from tqdm import tqdm
|
||||
|
||||
if len(sys.argv) != 2:
|
||||
print('You must enter the model name as a parameter, e.g.: download_model.py 124M')
|
||||
sys.exit(1)
|
||||
|
||||
model = sys.argv[1]
|
||||
|
||||
subdir = os.path.join('models', model)
|
||||
if not os.path.exists(subdir):
|
||||
os.makedirs(subdir)
|
||||
subdir = subdir.replace('\\','/') # needed for Windows
|
||||
|
||||
for filename in ['checkpoint','encoder.json','hparams.json','model.ckpt.data-00000-of-00001', 'model.ckpt.index', 'model.ckpt.meta', 'vocab.bpe']:
|
||||
|
||||
r = requests.get("https://storage.googleapis.com/gpt-2/" + subdir + "/" + filename, stream=True)
|
||||
|
||||
with open(os.path.join(subdir, filename), 'wb') as f:
|
||||
file_size = int(r.headers["content-length"])
|
||||
chunk_size = 1000
|
||||
with tqdm(ncols=100, desc="Fetching " + filename, total=file_size, unit_scale=True) as pbar:
|
||||
# 1k for chunk_size, since Ethernet packet size is around 1500 bytes
|
||||
for chunk in r.iter_content(chunk_size=chunk_size):
|
||||
f.write(chunk)
|
||||
pbar.update(chunk_size)
|
||||
@@ -0,0 +1,118 @@
|
||||
from story.utils import *
|
||||
import warnings
|
||||
warnings.filterwarnings("ignore")
|
||||
import os
|
||||
import requests
|
||||
import sys
|
||||
import tensorflow as tf
|
||||
import requests
|
||||
from tqdm import tqdm
|
||||
from generator.gpt2.src import model, sample, encoder
|
||||
import json
|
||||
import numpy as np
|
||||
|
||||
tf.logging.set_verbosity(tf.logging.ERROR)
|
||||
|
||||
class GPT2Generator:
|
||||
|
||||
def __init__(self, generate_num=80, temperature=0.3, top_k=40, top_p=0.8):
|
||||
self.generate_num=generate_num
|
||||
self.temp = temperature
|
||||
self.top_k = top_k
|
||||
self.top_p = top_p
|
||||
|
||||
self.model_name = "model_v1"
|
||||
self.model_dir = "generator/gpt2/models"
|
||||
self.checkpoint_path = os.path.join(self.model_dir, self.model_name)
|
||||
|
||||
models_dir = os.path.expanduser(os.path.expandvars(self.models_dir))
|
||||
self.batch_size = 1
|
||||
self.samples = 1
|
||||
|
||||
self.enc = encoder.get_encoder(self.model_name, models_dir)
|
||||
hparams = model.default_hparams()
|
||||
with open(os.path.join(models_dir, self.model_name, 'hparams.json')) as f:
|
||||
hparams.override_from_dict(json.load(f))
|
||||
seed = 20
|
||||
|
||||
self.sess = tf.Session(graph=tf.Graph())
|
||||
context = tf.placeholder(tf.int32, [self.batch_size, None])
|
||||
np.random.seed(seed)
|
||||
tf.set_random_seed(seed)
|
||||
self.output = sample.sample_sequence(
|
||||
hparams=hparams, length=self.generate_num,
|
||||
context=context,
|
||||
batch_size=self.batch_size,
|
||||
temperature=temperature, top_k=top_k, top_p=top_p
|
||||
)
|
||||
|
||||
saver = tf.train.Saver()
|
||||
ckpt = tf.train.latest_checkpoint(os.path.join(models_dir, self.model_name))
|
||||
saver.restore(self.sess, ckpt)
|
||||
|
||||
def prompt_replace(self, prompt):
|
||||
# print("\n\nBEFORE PROMPT_REPLACE:")
|
||||
# print(repr(prompt))
|
||||
if len(prompt) > 0 and prompt[-1] == " ":
|
||||
prompt = prompt[:-1]
|
||||
|
||||
#prompt = second_to_first_person(prompt)
|
||||
|
||||
# print("\n\nAFTER PROMPT_REPLACE")
|
||||
# print(repr(prompt))
|
||||
return prompt
|
||||
|
||||
def result_replace(self, result):
|
||||
# print("\n\nBEFORE RESULT_REPLACE:")
|
||||
# print(repr(result))
|
||||
|
||||
result = cut_trailing_sentence(result)
|
||||
if len(result) == 0:
|
||||
return ""
|
||||
first_letter_capitalized = result[0].isupper()
|
||||
result = result.replace('."', '".')
|
||||
result = result.replace("#", "")
|
||||
result = result.replace("*", "")
|
||||
#result = first_to_second_person(result)
|
||||
result = remove_profanity(result)
|
||||
|
||||
if not first_letter_capitalized:
|
||||
result = result[0].lower() + result[1:]
|
||||
|
||||
#
|
||||
# print("\n\nAFTER RESULT_REPLACE:")
|
||||
# print(repr(result))
|
||||
|
||||
return result
|
||||
|
||||
def generate(self, prompt, options=None, seed=1):
|
||||
|
||||
debug_print=False
|
||||
prefix = self.prompt_replace(prompt)
|
||||
|
||||
if debug_print:
|
||||
print("******DEBUG******")
|
||||
print("Prompt is: ", repr(prefix))
|
||||
|
||||
raw_text = input("Model prompt >>> ")
|
||||
|
||||
context_tokens = self.enc.encode(raw_text)
|
||||
generated = 0
|
||||
for _ in range(self.samples // self.batch_size):
|
||||
out = self.sess.run(self.output, feed_dict={
|
||||
self.context: [context_tokens for _ in range(self.batch_size)]
|
||||
})[:, len(context_tokens):]
|
||||
for i in range(self.batch_size):
|
||||
generated += 1
|
||||
text = self.enc.decode(out[i])
|
||||
|
||||
if debug_print:
|
||||
print("Generated result is: ", repr(text))
|
||||
print("******END DEBUG******")
|
||||
|
||||
result = text
|
||||
result = self.result_replace(result)
|
||||
if len(result) == 0:
|
||||
return self.generate(prompt)
|
||||
|
||||
return result
|
||||
@@ -0,0 +1,4 @@
|
||||
fire>=0.1.3
|
||||
regex==2018.1.10
|
||||
requests==2.21.0
|
||||
tqdm==4.31.1
|
||||
+11675
File diff suppressed because it is too large
Load Diff
Executable
+86
@@ -0,0 +1,86 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
import fire
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import numpy as np
|
||||
import tensorflow as tf
|
||||
|
||||
import model, sample, encoder
|
||||
|
||||
def interact_model(
|
||||
model_name='117M',
|
||||
seed=None,
|
||||
length=20,
|
||||
temperature=1,
|
||||
top_k=0,
|
||||
conversation="""
|
||||
you: hi
|
||||
her: hey
|
||||
you: i'm a human
|
||||
her: i'm a robot
|
||||
you: you ready?
|
||||
her: yes :)
|
||||
you: ok let's start chatting
|
||||
her: sure, what do you want to talk about?"""
|
||||
):
|
||||
|
||||
enc = encoder.get_encoder(model_name)
|
||||
hparams = model.default_hparams()
|
||||
with open(os.path.join('models', model_name, 'hparams.json')) as f:
|
||||
hparams.override_from_dict(json.load(f))
|
||||
|
||||
if length > hparams.n_ctx:
|
||||
raise ValueError("Can't get samples longer than window size: %s" % hparams.n_ctx)
|
||||
|
||||
with tf.Session(graph=tf.Graph()) as sess:
|
||||
np.random.seed(seed)
|
||||
tf.set_random_seed(seed)
|
||||
context = tf.placeholder(tf.int32, [1, None])
|
||||
output = sample.sample_sequence(
|
||||
hparams=hparams, length=length,
|
||||
context=context,
|
||||
batch_size=1,
|
||||
temperature=temperature, top_k=top_k
|
||||
)
|
||||
|
||||
print(conversation)
|
||||
|
||||
while True:
|
||||
saver = tf.train.Saver()
|
||||
ckpt = tf.train.latest_checkpoint(os.path.join('models', model_name))
|
||||
saver.restore(sess, ckpt)
|
||||
message = None
|
||||
while not message:
|
||||
message = input("you: ")
|
||||
conversation = conversation + "\nyou: " + message
|
||||
conversation = conversation + "\nher: "
|
||||
sys.stdout.write("her: ")
|
||||
sys.stdout.flush()
|
||||
|
||||
#sys.stderr.write("************************"+conversation+"***********************")
|
||||
#sys.stderr.flush()
|
||||
|
||||
encoded_conversation = enc.encode(conversation)
|
||||
#print(len(encoded_conversation))
|
||||
result = sess.run(output, feed_dict={
|
||||
context: [encoded_conversation]
|
||||
})[:, len(encoded_conversation):]
|
||||
text = enc.decode(result[0])
|
||||
|
||||
#sys.stderr.write("=============="+text+"=================")
|
||||
#sys.stderr.flush()
|
||||
|
||||
splits = text.split('\n')
|
||||
#line = splits[1] if len(splits)>1 else splits[0]
|
||||
#parts = line.split(': ')
|
||||
#reply = parts[1] if len(parts)>1 else parts[0]
|
||||
reply = splits[0]
|
||||
sys.stdout.write(reply+'\n')
|
||||
sys.stdout.flush()
|
||||
conversation = conversation + reply
|
||||
|
||||
if __name__ == '__main__':
|
||||
fire.Fire(interact_model)
|
||||
|
||||
@@ -0,0 +1,117 @@
|
||||
"""Byte pair encoding utilities"""
|
||||
|
||||
import os
|
||||
import json
|
||||
import regex as re
|
||||
from functools import lru_cache
|
||||
|
||||
@lru_cache()
|
||||
def bytes_to_unicode():
|
||||
"""
|
||||
Returns list of utf-8 byte and a corresponding list of unicode strings.
|
||||
The reversible bpe codes work on unicode strings.
|
||||
This means you need a large # of unicode characters in your vocab if you want to avoid UNKs.
|
||||
When you're at something like a 10B token dataset you end up needing around 5K for decent coverage.
|
||||
This is a signficant percentage of your normal, say, 32K bpe vocab.
|
||||
To avoid that, we want lookup tables between utf-8 bytes and unicode strings.
|
||||
And avoids mapping to whitespace/control characters the bpe code barfs on.
|
||||
"""
|
||||
bs = list(range(ord("!"), ord("~")+1))+list(range(ord("¡"), ord("¬")+1))+list(range(ord("®"), ord("ÿ")+1))
|
||||
cs = bs[:]
|
||||
n = 0
|
||||
for b in range(2**8):
|
||||
if b not in bs:
|
||||
bs.append(b)
|
||||
cs.append(2**8+n)
|
||||
n += 1
|
||||
cs = [chr(n) for n in cs]
|
||||
return dict(zip(bs, cs))
|
||||
|
||||
def get_pairs(word):
|
||||
"""Return set of symbol pairs in a word.
|
||||
|
||||
Word is represented as tuple of symbols (symbols being variable-length strings).
|
||||
"""
|
||||
pairs = set()
|
||||
prev_char = word[0]
|
||||
for char in word[1:]:
|
||||
pairs.add((prev_char, char))
|
||||
prev_char = char
|
||||
return pairs
|
||||
|
||||
class Encoder:
|
||||
def __init__(self, encoder, bpe_merges, errors='replace'):
|
||||
self.encoder = encoder
|
||||
self.decoder = {v:k for k,v in self.encoder.items()}
|
||||
self.errors = errors # how to handle errors in decoding
|
||||
self.byte_encoder = bytes_to_unicode()
|
||||
self.byte_decoder = {v:k for k, v in self.byte_encoder.items()}
|
||||
self.bpe_ranks = dict(zip(bpe_merges, range(len(bpe_merges))))
|
||||
self.cache = {}
|
||||
|
||||
# Should haved added re.IGNORECASE so BPE merges can happen for capitalized versions of contractions
|
||||
self.pat = re.compile(r"""'s|'t|'re|'ve|'m|'ll|'d| ?\p{L}+| ?\p{N}+| ?[^\s\p{L}\p{N}]+|\s+(?!\S)|\s+""")
|
||||
|
||||
def bpe(self, token):
|
||||
if token in self.cache:
|
||||
return self.cache[token]
|
||||
word = tuple(token)
|
||||
pairs = get_pairs(word)
|
||||
|
||||
if not pairs:
|
||||
return token
|
||||
|
||||
while True:
|
||||
bigram = min(pairs, key = lambda pair: self.bpe_ranks.get(pair, float('inf')))
|
||||
if bigram not in self.bpe_ranks:
|
||||
break
|
||||
first, second = bigram
|
||||
new_word = []
|
||||
i = 0
|
||||
while i < len(word):
|
||||
try:
|
||||
j = word.index(first, i)
|
||||
new_word.extend(word[i:j])
|
||||
i = j
|
||||
except:
|
||||
new_word.extend(word[i:])
|
||||
break
|
||||
|
||||
if word[i] == first and i < len(word)-1 and word[i+1] == second:
|
||||
new_word.append(first+second)
|
||||
i += 2
|
||||
else:
|
||||
new_word.append(word[i])
|
||||
i += 1
|
||||
new_word = tuple(new_word)
|
||||
word = new_word
|
||||
if len(word) == 1:
|
||||
break
|
||||
else:
|
||||
pairs = get_pairs(word)
|
||||
word = ' '.join(word)
|
||||
self.cache[token] = word
|
||||
return word
|
||||
|
||||
def encode(self, text):
|
||||
bpe_tokens = []
|
||||
for token in re.findall(self.pat, text):
|
||||
token = ''.join(self.byte_encoder[b] for b in token.encode('utf-8'))
|
||||
bpe_tokens.extend(self.encoder[bpe_token] for bpe_token in self.bpe(token).split(' '))
|
||||
return bpe_tokens
|
||||
|
||||
def decode(self, tokens):
|
||||
text = ''.join([self.decoder[token] for token in tokens])
|
||||
text = bytearray([self.byte_decoder[c] for c in text]).decode('utf-8', errors=self.errors)
|
||||
return text
|
||||
|
||||
def get_encoder(model_name, models_dir):
|
||||
with open(os.path.join(models_dir, model_name, 'encoder.json'), 'r') as f:
|
||||
encoder = json.load(f)
|
||||
with open(os.path.join(models_dir, model_name, 'vocab.bpe'), 'r', encoding="utf-8") as f:
|
||||
bpe_data = f.read()
|
||||
bpe_merges = [tuple(merge_str.split()) for merge_str in bpe_data.split('\n')[1:-1]]
|
||||
return Encoder(
|
||||
encoder=encoder,
|
||||
bpe_merges=bpe_merges,
|
||||
)
|
||||
+81
@@ -0,0 +1,81 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
import fire
|
||||
import json
|
||||
import os
|
||||
import numpy as np
|
||||
import tensorflow as tf
|
||||
|
||||
import model, sample, encoder
|
||||
|
||||
def sample_model(
|
||||
model_name='124M',
|
||||
seed=None,
|
||||
nsamples=0,
|
||||
batch_size=1,
|
||||
length=None,
|
||||
temperature=1,
|
||||
top_k=0,
|
||||
top_p=1,
|
||||
models_dir='models',
|
||||
):
|
||||
"""
|
||||
Run the sample_model
|
||||
:model_name=124M : String, which model to use
|
||||
:seed=None : Integer seed for random number generators, fix seed to
|
||||
reproduce results
|
||||
:nsamples=0 : Number of samples to return, if 0, continues to
|
||||
generate samples indefinately.
|
||||
:batch_size=1 : Number of batches (only affects speed/memory).
|
||||
:length=None : Number of tokens in generated text, if None (default), is
|
||||
determined by model hyperparameters
|
||||
:temperature=1 : Float value controlling randomness in boltzmann
|
||||
distribution. Lower temperature results in less random completions. As the
|
||||
temperature approaches zero, the model will become deterministic and
|
||||
repetitive. Higher temperature results in more random completions.
|
||||
:top_k=0 : Integer value controlling diversity. 1 means only 1 word is
|
||||
considered for each step (token), resulting in deterministic completions,
|
||||
while 40 means 40 words are considered at each step. 0 (default) is a
|
||||
special setting meaning no restrictions. 40 generally is a good value.
|
||||
:models_dir : path to parent folder containing model subfolders
|
||||
(i.e. contains the <model_name> folder)
|
||||
"""
|
||||
models_dir = os.path.expanduser(os.path.expandvars(models_dir))
|
||||
enc = encoder.get_encoder(model_name, models_dir)
|
||||
hparams = model.default_hparams()
|
||||
with open(os.path.join(models_dir, model_name, 'hparams.json')) as f:
|
||||
hparams.override_from_dict(json.load(f))
|
||||
|
||||
if length is None:
|
||||
length = hparams.n_ctx
|
||||
elif length > hparams.n_ctx:
|
||||
raise ValueError("Can't get samples longer than window size: %s" % hparams.n_ctx)
|
||||
|
||||
with tf.Session(graph=tf.Graph()) as sess:
|
||||
|
||||
np.random.seed(seed)
|
||||
tf.set_random_seed(seed)
|
||||
|
||||
output = sample.sample_sequence(
|
||||
hparams=hparams, length=length,
|
||||
start_token=enc.encoder['<|endoftext|>'],
|
||||
batch_size=batch_size,
|
||||
temperature=temperature, top_k=top_k, top_p=top_p
|
||||
)[:, 1:]
|
||||
|
||||
saver = tf.train.Saver()
|
||||
ckpt = tf.train.latest_checkpoint(os.path.join(models_dir, model_name))
|
||||
saver.restore(sess, ckpt)
|
||||
|
||||
generated = 0
|
||||
while nsamples == 0 or generated < nsamples:
|
||||
out = sess.run(output)
|
||||
for i in range(batch_size):
|
||||
generated += batch_size
|
||||
text = enc.decode(out[i])
|
||||
print("=" * 40 + " SAMPLE " + str(generated) + " " + "=" * 40)
|
||||
print(text)
|
||||
|
||||
if __name__ == '__main__':
|
||||
fire.Fire(sample_model)
|
||||
|
||||
+92
@@ -0,0 +1,92 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
import fire
|
||||
import json
|
||||
import os
|
||||
import numpy as np
|
||||
import tensorflow as tf
|
||||
|
||||
import model, sample, encoder
|
||||
|
||||
def interact_model(
|
||||
model_name='124M',
|
||||
seed=None,
|
||||
nsamples=1,
|
||||
batch_size=1,
|
||||
length=None,
|
||||
temperature=1,
|
||||
top_k=0,
|
||||
top_p=1,
|
||||
models_dir='models',
|
||||
):
|
||||
"""
|
||||
Interactively run the model
|
||||
:model_name=124M : String, which model to use
|
||||
:seed=None : Integer seed for random number generators, fix seed to reproduce
|
||||
results
|
||||
:nsamples=1 : Number of samples to return total
|
||||
:batch_size=1 : Number of batches (only affects speed/memory). Must divide nsamples.
|
||||
:length=None : Number of tokens in generated text, if None (default), is
|
||||
determined by model hyperparameters
|
||||
:temperature=1 : Float value controlling randomness in boltzmann
|
||||
distribution. Lower temperature results in less random completions. As the
|
||||
temperature approaches zero, the model will become deterministic and
|
||||
repetitive. Higher temperature results in more random completions.
|
||||
:top_k=0 : Integer value controlling diversity. 1 means only 1 word is
|
||||
considered for each step (token), resulting in deterministic completions,
|
||||
while 40 means 40 words are considered at each step. 0 (default) is a
|
||||
special setting meaning no restrictions. 40 generally is a good value.
|
||||
:models_dir : path to parent folder containing model subfolders
|
||||
(i.e. contains the <model_name> folder)
|
||||
"""
|
||||
models_dir = os.path.expanduser(os.path.expandvars(models_dir))
|
||||
if batch_size is None:
|
||||
batch_size = 1
|
||||
assert nsamples % batch_size == 0
|
||||
|
||||
enc = encoder.get_encoder(model_name, models_dir)
|
||||
hparams = model.default_hparams()
|
||||
with open(os.path.join(models_dir, model_name, 'hparams.json')) as f:
|
||||
hparams.override_from_dict(json.load(f))
|
||||
|
||||
if length is None:
|
||||
length = hparams.n_ctx // 2
|
||||
elif length > hparams.n_ctx:
|
||||
raise ValueError("Can't get samples longer than window size: %s" % hparams.n_ctx)
|
||||
|
||||
with tf.Session(graph=tf.Graph()) as sess:
|
||||
context = tf.placeholder(tf.int32, [batch_size, None])
|
||||
np.random.seed(seed)
|
||||
tf.set_random_seed(seed)
|
||||
output = sample.sample_sequence(
|
||||
hparams=hparams, length=length,
|
||||
context=context,
|
||||
batch_size=batch_size,
|
||||
temperature=temperature, top_k=top_k, top_p=top_p
|
||||
)
|
||||
|
||||
saver = tf.train.Saver()
|
||||
ckpt = tf.train.latest_checkpoint(os.path.join(models_dir, model_name))
|
||||
saver.restore(sess, ckpt)
|
||||
|
||||
while True:
|
||||
raw_text = input("Model prompt >>> ")
|
||||
while not raw_text:
|
||||
print('Prompt should not be empty!')
|
||||
raw_text = raw_input("Model prompt >>> ")
|
||||
context_tokens = enc.encode(raw_text)
|
||||
generated = 0
|
||||
for _ in range(nsamples // batch_size):
|
||||
out = sess.run(output, feed_dict={
|
||||
context: [context_tokens for _ in range(batch_size)]
|
||||
})[:, len(context_tokens):]
|
||||
for i in range(batch_size):
|
||||
generated += 1
|
||||
text = enc.decode(out[i])
|
||||
print("=" * 40 + " SAMPLE " + str(generated) + " " + "=" * 40)
|
||||
print(text)
|
||||
print("=" * 80)
|
||||
|
||||
if __name__ == '__main__':
|
||||
fire.Fire(interact_model)
|
||||
|
||||
@@ -0,0 +1,174 @@
|
||||
import numpy as np
|
||||
import tensorflow as tf
|
||||
from tensorflow.contrib.training import HParams
|
||||
|
||||
def default_hparams():
|
||||
return HParams(
|
||||
n_vocab=0,
|
||||
n_ctx=1024,
|
||||
n_embd=768,
|
||||
n_head=12,
|
||||
n_layer=12,
|
||||
)
|
||||
|
||||
def shape_list(x):
|
||||
"""Deal with dynamic shape in tensorflow cleanly."""
|
||||
static = x.shape.as_list()
|
||||
dynamic = tf.shape(x)
|
||||
return [dynamic[i] if s is None else s for i, s in enumerate(static)]
|
||||
|
||||
def softmax(x, axis=-1):
|
||||
x = x - tf.reduce_max(x, axis=axis, keepdims=True)
|
||||
ex = tf.exp(x)
|
||||
return ex / tf.reduce_sum(ex, axis=axis, keepdims=True)
|
||||
|
||||
def gelu(x):
|
||||
return 0.5*x*(1+tf.tanh(np.sqrt(2/np.pi)*(x+0.044715*tf.pow(x, 3))))
|
||||
|
||||
def norm(x, scope, *, axis=-1, epsilon=1e-5):
|
||||
"""Normalize to mean = 0, std = 1, then do a diagonal affine transform."""
|
||||
with tf.variable_scope(scope):
|
||||
n_state = x.shape[-1].value
|
||||
g = tf.get_variable('g', [n_state], initializer=tf.constant_initializer(1))
|
||||
b = tf.get_variable('b', [n_state], initializer=tf.constant_initializer(0))
|
||||
u = tf.reduce_mean(x, axis=axis, keepdims=True)
|
||||
s = tf.reduce_mean(tf.square(x-u), axis=axis, keepdims=True)
|
||||
x = (x - u) * tf.rsqrt(s + epsilon)
|
||||
x = x*g + b
|
||||
return x
|
||||
|
||||
def split_states(x, n):
|
||||
"""Reshape the last dimension of x into [n, x.shape[-1]/n]."""
|
||||
*start, m = shape_list(x)
|
||||
return tf.reshape(x, start + [n, m//n])
|
||||
|
||||
def merge_states(x):
|
||||
"""Smash the last two dimensions of x into a single dimension."""
|
||||
*start, a, b = shape_list(x)
|
||||
return tf.reshape(x, start + [a*b])
|
||||
|
||||
def conv1d(x, scope, nf, *, w_init_stdev=0.02):
|
||||
with tf.variable_scope(scope):
|
||||
*start, nx = shape_list(x)
|
||||
w = tf.get_variable('w', [1, nx, nf], initializer=tf.random_normal_initializer(stddev=w_init_stdev))
|
||||
b = tf.get_variable('b', [nf], initializer=tf.constant_initializer(0))
|
||||
c = tf.reshape(tf.matmul(tf.reshape(x, [-1, nx]), tf.reshape(w, [-1, nf]))+b, start+[nf])
|
||||
return c
|
||||
|
||||
def attention_mask(nd, ns, *, dtype):
|
||||
"""1's in the lower triangle, counting from the lower right corner.
|
||||
|
||||
Same as tf.matrix_band_part(tf.ones([nd, ns]), -1, ns-nd), but doesn't produce garbage on TPUs.
|
||||
"""
|
||||
i = tf.range(nd)[:,None]
|
||||
j = tf.range(ns)
|
||||
m = i >= j - ns + nd
|
||||
return tf.cast(m, dtype)
|
||||
|
||||
|
||||
def attn(x, scope, n_state, *, past, hparams):
|
||||
assert x.shape.ndims == 3 # Should be [batch, sequence, features]
|
||||
assert n_state % hparams.n_head == 0
|
||||
if past is not None:
|
||||
assert past.shape.ndims == 5 # Should be [batch, 2, heads, sequence, features], where 2 is [k, v]
|
||||
|
||||
def split_heads(x):
|
||||
# From [batch, sequence, features] to [batch, heads, sequence, features]
|
||||
return tf.transpose(split_states(x, hparams.n_head), [0, 2, 1, 3])
|
||||
|
||||
def merge_heads(x):
|
||||
# Reverse of split_heads
|
||||
return merge_states(tf.transpose(x, [0, 2, 1, 3]))
|
||||
|
||||
def mask_attn_weights(w):
|
||||
# w has shape [batch, heads, dst_sequence, src_sequence], where information flows from src to dst.
|
||||
_, _, nd, ns = shape_list(w)
|
||||
b = attention_mask(nd, ns, dtype=w.dtype)
|
||||
b = tf.reshape(b, [1, 1, nd, ns])
|
||||
w = w*b - tf.cast(1e10, w.dtype)*(1-b)
|
||||
return w
|
||||
|
||||
def multihead_attn(q, k, v):
|
||||
# q, k, v have shape [batch, heads, sequence, features]
|
||||
w = tf.matmul(q, k, transpose_b=True)
|
||||
w = w * tf.rsqrt(tf.cast(v.shape[-1].value, w.dtype))
|
||||
|
||||
w = mask_attn_weights(w)
|
||||
w = softmax(w)
|
||||
a = tf.matmul(w, v)
|
||||
return a
|
||||
|
||||
with tf.variable_scope(scope):
|
||||
c = conv1d(x, 'c_attn', n_state*3)
|
||||
q, k, v = map(split_heads, tf.split(c, 3, axis=2))
|
||||
present = tf.stack([k, v], axis=1)
|
||||
if past is not None:
|
||||
pk, pv = tf.unstack(past, axis=1)
|
||||
k = tf.concat([pk, k], axis=-2)
|
||||
v = tf.concat([pv, v], axis=-2)
|
||||
a = multihead_attn(q, k, v)
|
||||
a = merge_heads(a)
|
||||
a = conv1d(a, 'c_proj', n_state)
|
||||
return a, present
|
||||
|
||||
|
||||
def mlp(x, scope, n_state, *, hparams):
|
||||
with tf.variable_scope(scope):
|
||||
nx = x.shape[-1].value
|
||||
h = gelu(conv1d(x, 'c_fc', n_state))
|
||||
h2 = conv1d(h, 'c_proj', nx)
|
||||
return h2
|
||||
|
||||
|
||||
def block(x, scope, *, past, hparams):
|
||||
with tf.variable_scope(scope):
|
||||
nx = x.shape[-1].value
|
||||
a, present = attn(norm(x, 'ln_1'), 'attn', nx, past=past, hparams=hparams)
|
||||
x = x + a
|
||||
m = mlp(norm(x, 'ln_2'), 'mlp', nx*4, hparams=hparams)
|
||||
x = x + m
|
||||
return x, present
|
||||
|
||||
def past_shape(*, hparams, batch_size=None, sequence=None):
|
||||
return [batch_size, hparams.n_layer, 2, hparams.n_head, sequence, hparams.n_embd // hparams.n_head]
|
||||
|
||||
def expand_tile(value, size):
|
||||
"""Add a new axis of given size."""
|
||||
value = tf.convert_to_tensor(value, name='value')
|
||||
ndims = value.shape.ndims
|
||||
return tf.tile(tf.expand_dims(value, axis=0), [size] + [1]*ndims)
|
||||
|
||||
def positions_for(tokens, past_length):
|
||||
batch_size = tf.shape(tokens)[0]
|
||||
nsteps = tf.shape(tokens)[1]
|
||||
return expand_tile(past_length + tf.range(nsteps), batch_size)
|
||||
|
||||
|
||||
def model(hparams, X, past=None, scope='model', reuse=False):
|
||||
with tf.variable_scope(scope, reuse=reuse):
|
||||
results = {}
|
||||
batch, sequence = shape_list(X)
|
||||
|
||||
wpe = tf.get_variable('wpe', [hparams.n_ctx, hparams.n_embd],
|
||||
initializer=tf.random_normal_initializer(stddev=0.01))
|
||||
wte = tf.get_variable('wte', [hparams.n_vocab, hparams.n_embd],
|
||||
initializer=tf.random_normal_initializer(stddev=0.02))
|
||||
past_length = 0 if past is None else tf.shape(past)[-2]
|
||||
h = tf.gather(wte, X) + tf.gather(wpe, positions_for(X, past_length))
|
||||
|
||||
# Transformer
|
||||
presents = []
|
||||
pasts = tf.unstack(past, axis=1) if past is not None else [None] * hparams.n_layer
|
||||
assert len(pasts) == hparams.n_layer
|
||||
for layer, past in enumerate(pasts):
|
||||
h, present = block(h, 'h%d' % layer, past=past, hparams=hparams)
|
||||
presents.append(present)
|
||||
results['present'] = tf.stack(presents, axis=1)
|
||||
h = norm(h, 'ln_f')
|
||||
|
||||
# Language model loss. Do tokens <n predict token n?
|
||||
h_flat = tf.reshape(h, [batch*sequence, hparams.n_embd])
|
||||
logits = tf.matmul(h_flat, wte, transpose_b=True)
|
||||
logits = tf.reshape(logits, [batch, sequence, hparams.n_vocab])
|
||||
results['logits'] = logits
|
||||
return results
|
||||
@@ -0,0 +1,114 @@
|
||||
import tensorflow as tf
|
||||
|
||||
import model
|
||||
|
||||
def penalize_used(logits, output):
|
||||
|
||||
# I want to change the indices of logits wherever the index is found in output
|
||||
change_tensor = tf.zeros_like(logits, dtype=logits.dtype)
|
||||
unique = tf.unique(output[0])[0]
|
||||
ones = tf.ones_like(unique, dtype=unique.dtype)
|
||||
indices = tf.expand_dims(unique, 1)
|
||||
|
||||
updates = tf.scatter_nd(indices, ones, [logits.shape[1]])
|
||||
|
||||
bool_tensor = tf.expand_dims(tf.cast(updates, tf.bool), 0)
|
||||
|
||||
return tf.compat.v1.where(
|
||||
bool_tensor,
|
||||
logits / 1.2,
|
||||
logits)
|
||||
|
||||
|
||||
def top_k_logits(logits, k):
|
||||
if k == 0:
|
||||
# no truncation
|
||||
return logits
|
||||
|
||||
def _top_k():
|
||||
values, _ = tf.nn.top_k(logits, k=k)
|
||||
min_values = values[:, -1, tf.newaxis]
|
||||
return tf.where(
|
||||
logits < min_values,
|
||||
tf.ones_like(logits, dtype=logits.dtype) * -1e10,
|
||||
logits,
|
||||
)
|
||||
return tf.cond(
|
||||
tf.equal(k, 0),
|
||||
lambda: logits,
|
||||
lambda: _top_k(),
|
||||
)
|
||||
|
||||
|
||||
def top_p_logits(logits, p):
|
||||
"""Nucleus sampling"""
|
||||
batch, _ = logits.shape.as_list()
|
||||
sorted_logits = tf.sort(logits, direction='DESCENDING', axis=-1)
|
||||
cumulative_probs = tf.cumsum(tf.nn.softmax(sorted_logits, axis=-1), axis=-1)
|
||||
indices = tf.stack([
|
||||
tf.range(0, batch),
|
||||
# number of indices to include
|
||||
tf.maximum(tf.reduce_sum(tf.cast(cumulative_probs <= p, tf.int32), axis=-1) - 1, 0),
|
||||
], axis=-1)
|
||||
min_values = tf.gather_nd(sorted_logits, indices)
|
||||
return tf.where(
|
||||
logits < min_values,
|
||||
tf.ones_like(logits) * -1e10,
|
||||
logits,
|
||||
)
|
||||
|
||||
|
||||
def sample_sequence(*, hparams, length, start_token=None, batch_size=None, context=None, temperature=1, top_k=0, top_p=1):
|
||||
if start_token is None:
|
||||
assert context is not None, 'Specify exactly one of start_token and context!'
|
||||
else:
|
||||
assert context is None, 'Specify exactly one of start_token and context!'
|
||||
context = tf.fill([batch_size, 1], start_token)
|
||||
|
||||
def step(hparams, tokens, past=None):
|
||||
lm_output = model.model(hparams=hparams, X=tokens, past=past, reuse=tf.AUTO_REUSE)
|
||||
|
||||
logits = lm_output['logits'][:, :, :hparams.n_vocab]
|
||||
presents = lm_output['present']
|
||||
presents.set_shape(model.past_shape(hparams=hparams, batch_size=batch_size))
|
||||
return {
|
||||
'logits': logits,
|
||||
'presents': presents,
|
||||
}
|
||||
|
||||
with tf.name_scope('sample_sequence'):
|
||||
def body(past, prev, output):
|
||||
next_outputs = step(hparams, prev, past=past)
|
||||
logits = next_outputs['logits'][:, -1, :] / tf.to_float(temperature)
|
||||
logits = penalize_used(logits, output)
|
||||
logits = top_k_logits(logits, k=top_k)
|
||||
logits = top_p_logits(logits, p=top_p)
|
||||
samples = tf.multinomial(logits, num_samples=1, output_dtype=tf.int32)
|
||||
return [
|
||||
next_outputs['presents'] if past is None else tf.concat([past, next_outputs['presents']], axis=-2),
|
||||
samples,
|
||||
tf.concat([output, samples], axis=1)
|
||||
]
|
||||
|
||||
past, prev, output = body(None, context, context)
|
||||
|
||||
def cond(*args):
|
||||
return True
|
||||
|
||||
_, _, tokens = tf.while_loop(
|
||||
cond=cond, body=body,
|
||||
maximum_iterations=length - 1,
|
||||
loop_vars=[
|
||||
past,
|
||||
prev,
|
||||
output
|
||||
],
|
||||
shape_invariants=[
|
||||
tf.TensorShape(model.past_shape(hparams=hparams, batch_size=batch_size)),
|
||||
tf.TensorShape([batch_size, None]),
|
||||
tf.TensorShape([batch_size, None]),
|
||||
],
|
||||
back_prop=False,
|
||||
)
|
||||
|
||||
return tokens
|
||||
@@ -0,0 +1,15 @@
|
||||
import numpy as np
|
||||
import tensorflow as tf
|
||||
|
||||
|
||||
|
||||
seed=None
|
||||
|
||||
print(np.random.seed(seed))
|
||||
#
|
||||
|
||||
tf.set_random_seed(seed)
|
||||
generate = tf.random_uniform(())
|
||||
with tf.Session() as sess:
|
||||
print(generate.eval())
|
||||
# 0.96046877
|
||||
@@ -1,3 +1,4 @@
|
||||
models
|
||||
*.pyc
|
||||
__pycache__
|
||||
__pycache__
|
||||
checkpoint
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,147 +0,0 @@
|
||||
from story.utils import *
|
||||
import warnings
|
||||
warnings.filterwarnings("ignore")
|
||||
import gpt_2_simple as gpt2
|
||||
import os
|
||||
import requests
|
||||
import sys
|
||||
import tensorflow as tf
|
||||
import requests
|
||||
from tqdm import tqdm
|
||||
from gpt_2_simple.src import model, sample, encoder, memory_saving_gradients
|
||||
from gpt_2_simple.src.load_dataset import load_dataset, Sampler
|
||||
from gpt_2_simple.src.accumulate import AccumulatingOptimizer
|
||||
import json
|
||||
import numpy as np
|
||||
|
||||
tf.logging.set_verbosity(tf.logging.ERROR)
|
||||
|
||||
class SimpleGenerator:
|
||||
|
||||
def __init__(self, generate_num=80, temperature=0.3, top_k=40, top_p=0.8):
|
||||
self.generate_num=generate_num
|
||||
self.temp = temperature
|
||||
self.top_k = top_k
|
||||
self.top_p = top_p
|
||||
|
||||
self.model_name = "run1"
|
||||
self.model_dir = "generator/simple/checkpoint"
|
||||
self.checkpoint_path = os.path.join(self.model_dir, self.model_name)
|
||||
if not os.path.isdir(os.path.join(self.model_dir, self.model_name)):
|
||||
print(f"Downloading {self.model_name} model...")
|
||||
|
||||
subdir = os.path.join(self.model_dir, self.model_name)
|
||||
if not os.path.exists(subdir):
|
||||
os.makedirs(subdir)
|
||||
subdir = subdir.replace('\\', '/') # needed for Windows
|
||||
|
||||
for filename in ['checkpoint', 'encoder.json', 'hparams.json', 'model.ckpt.data-00000-of-00001',
|
||||
'model.ckpt.index', 'model.ckpt.meta', 'vocab.bpe']:
|
||||
|
||||
r = requests.get("https://storage.googleapis.com/gpt-2/" + subdir + "/" + filename, stream=True)
|
||||
|
||||
with open(os.path.join(subdir, filename), 'wb') as f:
|
||||
file_size = int(r.headers["content-length"])
|
||||
chunk_size = 1000
|
||||
with tqdm(ncols=100, desc="Fetching " + filename, total=file_size, unit_scale=True) as pbar:
|
||||
# 1k for chunk_size, since Ethernet packet size is around 1500 bytes
|
||||
for chunk in r.iter_content(chunk_size=chunk_size):
|
||||
f.write(chunk)
|
||||
pbar.update(chunk_size)
|
||||
|
||||
self.sess = gpt2.start_tf_sess()
|
||||
gpt2.load_gpt2(self.sess, model_dir=self.model_dir, model_name=self.model_name)
|
||||
|
||||
|
||||
def prompt_replace(self, prompt):
|
||||
# print("\n\nBEFORE PROMPT_REPLACE:")
|
||||
# print(repr(prompt))
|
||||
if len(prompt) > 0 and prompt[-1] == " ":
|
||||
prompt = prompt[:-1]
|
||||
|
||||
#prompt = second_to_first_person(prompt)
|
||||
|
||||
# print("\n\nAFTER PROMPT_REPLACE")
|
||||
# print(repr(prompt))
|
||||
return prompt
|
||||
|
||||
def result_replace(self, result):
|
||||
# print("\n\nBEFORE RESULT_REPLACE:")
|
||||
# print(repr(result))
|
||||
|
||||
result = cut_trailing_sentence(result)
|
||||
if len(result) == 0:
|
||||
return ""
|
||||
first_letter_capitalized = result[0].isupper()
|
||||
result = result.replace('."', '".')
|
||||
result = result.replace("#", "")
|
||||
result = result.replace("*", "")
|
||||
#result = first_to_second_person(result)
|
||||
result = remove_profanity(result)
|
||||
|
||||
if not first_letter_capitalized:
|
||||
result = result[0].lower() + result[1:]
|
||||
|
||||
#
|
||||
# print("\n\nAFTER RESULT_REPLACE:")
|
||||
# print(repr(result))
|
||||
|
||||
return result
|
||||
|
||||
def generate(self, prompt, options=None, seed=1):
|
||||
|
||||
debug_print=False
|
||||
prefix = self.prompt_replace(prompt)
|
||||
|
||||
if debug_print:
|
||||
print("******DEBUG******")
|
||||
print("Prompt is: ", repr(prefix))
|
||||
|
||||
|
||||
enc = encoder.get_encoder(self.checkpoint_path)
|
||||
hparams = model.default_hparams()
|
||||
batch_size = 1
|
||||
nsamples = 1
|
||||
with open(os.path.join(self.checkpoint_path, 'hparams.json')) as f:
|
||||
hparams.override_from_dict(json.load(f))
|
||||
|
||||
context = tf.compat.v1.placeholder(tf.int32, [batch_size, None])
|
||||
context_tokens = enc.encode(prefix)
|
||||
|
||||
#np.random.seed(seed)
|
||||
#tf.compat.v1.set_random_seed(seed)
|
||||
|
||||
output = sample.sample_sequence(
|
||||
hparams=hparams,
|
||||
length=min(self.generate_num, 1023 - (len(context_tokens) if prefix else 0)),
|
||||
start_token=enc.encoder['<|endoftext|>'] if not prefix else None,
|
||||
context=context if prefix else None,
|
||||
batch_size=batch_size,
|
||||
temperature=self.temp, top_p=self.top_p
|
||||
)[:, 1:]
|
||||
|
||||
generated = 0
|
||||
gen_texts = []
|
||||
while generated < nsamples:
|
||||
if not prefix:
|
||||
out = self.sess.run(output)
|
||||
else:
|
||||
out = self.sess.run(output, feed_dict={
|
||||
context: batch_size * [context_tokens]
|
||||
})
|
||||
for i in range(batch_size):
|
||||
generated += 1
|
||||
gen_text = enc.decode(out[i])
|
||||
if prefix:
|
||||
gen_text = enc.decode(context_tokens[:1]) + gen_text
|
||||
gen_text = gen_text.lstrip('\n')
|
||||
gen_texts.append(gen_text)
|
||||
if debug_print:
|
||||
print("Generated result is: ", repr(gen_texts[0]))
|
||||
print("******END DEBUG******")
|
||||
result = gen_texts[0][len(prefix):]
|
||||
|
||||
result = self.result_replace(result)
|
||||
if len(result) == 0:
|
||||
return self.generate(prompt)
|
||||
return result
|
||||
@@ -1,14 +0,0 @@
|
||||
import gpt_2_simple as gpt2
|
||||
import os
|
||||
import requests
|
||||
|
||||
model_name = "1558M"
|
||||
if not os.path.isdir(os.path.join("models", model_name)):
|
||||
print(f"Downloading {model_name} model...")
|
||||
gpt2.download_gpt2(model_name=model_name) # model is saved into current directory under /models/124M/
|
||||
|
||||
sess = gpt2.start_tf_sess()
|
||||
gpt2.load_gpt2(sess, model_name=model_name)
|
||||
gpt2.generate(sess, model_name=model_name, length=30, prefix="I wake up in an old rundown hospital. I look around and see")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user