From 6b401119f974f1ffaec035fd1e1ac7d39e9d137a Mon Sep 17 00:00:00 2001 From: sirmammingtonham Date: Fri, 12 Jul 2019 16:58:37 -0700 Subject: [PATCH 1/2] allows for non-fixed length generation --- gpt_2_simple/gpt_2.py | 116 ++++++++++++++++++++++++++---------------- 1 file changed, 73 insertions(+), 43 deletions(-) diff --git a/gpt_2_simple/gpt_2.py b/gpt_2_simple/gpt_2.py index e02021b..6a78d51 100644 --- a/gpt_2_simple/gpt_2.py +++ b/gpt_2_simple/gpt_2.py @@ -362,7 +362,9 @@ def generate(sess, temperature=0.7, top_k=0, top_p=0.0, - include_prefix=True): + include_prefix=True, + split_context=0.5, + ): """Generates text from a model loaded into memory. Adapted from https://github.com/openai/gpt-2/blob/master/src/interactive_conditional_samples.py @@ -378,6 +380,10 @@ def generate(sess, if prefix == '': prefix = None + if not length: + assert truncate is not None, "If generating a non-fixed length \ + sample, must have a truncation term." + CHECKPOINT_DIR = 'checkpoint' SAMPLE_DIR = 'samples' @@ -388,56 +394,74 @@ def generate(sess, with open(os.path.join(checkpoint_path, 'hparams.json')) as f: hparams.override_from_dict(json.load(f)) - if prefix: - context = tf.placeholder(tf.int32, [batch_size, None]) - context_tokens = enc.encode(prefix) - np.random.seed(seed) tf.set_random_seed(seed) - output = sample.sample_sequence( - hparams=hparams, - length=min(length, 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=temperature, top_k=top_k, top_p=top_p - )[:, 1:] - if destination_path: f = open(destination_path, 'w') generated = 0 gen_texts = [] while generated < nsamples: - if not prefix: - out = sess.run(output) - else: - out = sess.run(output, feed_dict={ - context: batch_size * [context_tokens] - }) - for i in range(batch_size): - generated += 1 - gen_text = enc.decode(out[i]) + trunc_text = '' + gen_text = [''] * batch_size + while not trunc_text: if prefix: - gen_text = enc.decode(context_tokens[:1]) + gen_text - if truncate: - truncate_esc = re.escape(truncate) - if prefix and not include_prefix: - prefix_esc = re.escape(prefix) - pattern = '(?:{})(.*?)(?:{})'.format(prefix_esc, - truncate_esc) + context = tf.placeholder(tf.int32, [batch_size, None]) + context_tokens = enc.encode(prefix) + output = sample.sample_sequence( + hparams=hparams, + length=min(length if length else 1023, 1023 - (len(context_tokens))), + start_token=None, + context=context, + batch_size=batch_size, + temperature=temperature, top_k=top_k, top_p=top_p + )[:, 1:] + out = sess.run(output, feed_dict={ + context: batch_size * [context_tokens] + }) + + else: + output = sample.sample_sequence( + hparams=hparams, + length=min(length if length else 1023, 1023), + start_token=enc.encoder['<|endoftext|>'], + context=None, + batch_size=batch_size, + temperature=temperature, top_k=top_k, top_p=top_p + )[:, 1:] + out = sess.run(output) + + for i in range(batch_size): + text = enc.decode(out[i]) + if prefix: + text = enc.decode(context_tokens[:1]) + text + if truncate and not length: + split = [s + ' ' for s in text.split(' ')] + prefix = ''.join(split[int(len(split)*split_context):]) + + truncate_esc = re.escape(truncate) + if prefix and not include_prefix: + prefix_esc = re.escape(prefix) + pattern = '(?:{})(.*?)(?:{})'.format(prefix_esc, + truncate_esc) + else: + pattern = '(.*?)(?:{})'.format(truncate_esc) + + trunc_text = re.search(pattern, text, re.S) + if trunc_text: + text = trunc_text.group(1) else: - pattern = '(.*?)(?:{})'.format(truncate_esc) + trunc_text = True + gen_text[i] += text.lstrip('\n') - trunc_text = re.search(pattern, gen_text, re.S) - if trunc_text: - gen_text = trunc_text.group(1) - gen_text = gen_text.lstrip('\n') + for gen in gen_text: if destination_path: - f.write("{}\n{}".format(gen_text, sample_delim)) + f.write("{}\n{}".format(gen, sample_delim)) if not return_as_list and not destination_path: - print("{}\n{}".format(gen_text, sample_delim), end='') - gen_texts.append(gen_text) + print("{}\n{}".format(gen, sample_delim), end='') + gen_texts.append(gen) + + generated += batch_size if destination_path: f.close() @@ -459,7 +483,8 @@ def generate_to_file(sess, temperature=0.7, top_k=0, top_p=0.0, - include_prefix=True): + include_prefix=True, + split_context=0.5): """Generates the texts to a file. sample_delim separates texts: set to '' if each text is a small document. @@ -481,7 +506,8 @@ def generate_to_file(sess, temperature, top_k, top_p, - include_prefix) + include_prefix, + split_context) def mount_gdrive(): @@ -671,6 +697,9 @@ def cmd(): parser.add_argument( '--sample_delim', help="[generate] Delimiter between each generated sample.", nargs='?', default='=' * 20 + '\n', type=str) + parser.add_argument( + '--split_context', help="[generate] When generating a sample of non-fixed length, feed this proportion of previous generation as context.", + nargs='?', default=0.5, type=float) # Positional arguments parser.add_argument('mode', nargs='?') @@ -696,7 +725,7 @@ def cmd(): prefix=args.prefix, truncate=args.truncate, include_prefix=args.include_prefix, sample_delim=args.sample_delim, run_name=args.run_name, - top_k=args.top_k, top_p=args.top_p) + top_k=args.top_k, top_p=args.top_p, split_context=args.split_context) def cmd_finetune(dataset, run_name, model_name, steps, @@ -720,7 +749,7 @@ def cmd_generate(nfiles, nsamples, folder, length, temperature, batch_size, prefix, truncate, include_prefix, sample_delim, run_name, - top_k, top_p): + top_k, top_p, split_context): """Wrapper script for generating text via the CLI. The files are generated into a folder, which can be downloaded recursively by downloading the entire folder. @@ -751,5 +780,6 @@ def cmd_generate(nfiles, nsamples, folder, include_prefix=include_prefix, sample_delim=sample_delim, top_k=top_k, - top_p=top_p + top_p=top_p, + split_context=split_context ) From 5053cf593e8738e04a0e3ae0d576868df4d611be Mon Sep 17 00:00:00 2001 From: sirmammingtonham Date: Fri, 19 Jul 2019 17:08:21 -0700 Subject: [PATCH 2/2] fixed the problems with prev commit --- gpt_2_simple/gpt_2.py | 102 +++++++++++++++++++++--------------------- 1 file changed, 51 insertions(+), 51 deletions(-) diff --git a/gpt_2_simple/gpt_2.py b/gpt_2_simple/gpt_2.py index 6a78d51..96699c1 100644 --- a/gpt_2_simple/gpt_2.py +++ b/gpt_2_simple/gpt_2.py @@ -363,8 +363,7 @@ def generate(sess, top_k=0, top_p=0.0, include_prefix=True, - split_context=0.5, - ): + split_context=0.5): """Generates text from a model loaded into memory. Adapted from https://github.com/openai/gpt-2/blob/master/src/interactive_conditional_samples.py @@ -393,7 +392,12 @@ def generate(sess, hparams = model.default_hparams() with open(os.path.join(checkpoint_path, 'hparams.json')) as f: hparams.override_from_dict(json.load(f)) - + + context = tf.placeholder(tf.int32, [batch_size, None]) + if prefix: + context_tokens = [enc.encode(prefix)] * batch_size + else: + context_tokens = [[enc.encoder['<|endoftext|>']] for _ in range(batch_size)] np.random.seed(seed) tf.set_random_seed(seed) @@ -402,57 +406,53 @@ def generate(sess, generated = 0 gen_texts = [] while generated < nsamples: - trunc_text = '' gen_text = [''] * batch_size - while not trunc_text: - if prefix: - context = tf.placeholder(tf.int32, [batch_size, None]) - context_tokens = enc.encode(prefix) - output = sample.sample_sequence( - hparams=hparams, - length=min(length if length else 1023, 1023 - (len(context_tokens))), - start_token=None, - context=context, - batch_size=batch_size, - temperature=temperature, top_k=top_k, top_p=top_p - )[:, 1:] - out = sess.run(output, feed_dict={ - context: batch_size * [context_tokens] - }) - - else: - output = sample.sample_sequence( - hparams=hparams, - length=min(length if length else 1023, 1023), - start_token=enc.encoder['<|endoftext|>'], - context=None, - batch_size=batch_size, - temperature=temperature, top_k=top_k, top_p=top_p - )[:, 1:] - out = sess.run(output) - + truncated = [False] * batch_size + total_tokens = 0 + + while False in truncated: + num_tokens = 1023 - (len(context_tokens[0])) + output = sample.sample_sequence( + hparams=hparams, + length=min(length if length else 1023, num_tokens), + context=context, + batch_size=batch_size, + temperature=temperature, top_k=top_k, top_p=top_p + )[:, 1:] + + out = sess.run(output, feed_dict={ + context: context_tokens + }) + + total_tokens += num_tokens + for i in range(batch_size): text = enc.decode(out[i]) if prefix: - text = enc.decode(context_tokens[:1]) + text - if truncate and not length: - split = [s + ' ' for s in text.split(' ')] - prefix = ''.join(split[int(len(split)*split_context):]) - - truncate_esc = re.escape(truncate) - if prefix and not include_prefix: - prefix_esc = re.escape(prefix) - pattern = '(?:{})(.*?)(?:{})'.format(prefix_esc, - truncate_esc) - else: - pattern = '(.*?)(?:{})'.format(truncate_esc) - - trunc_text = re.search(pattern, text, re.S) - if trunc_text: - text = trunc_text.group(1) - else: - trunc_text = True - gen_text[i] += text.lstrip('\n') + text = enc.decode(context_tokens[i][:1]) + text + if truncate or all(gen_text): + context_tokens[i] = out[i][int(len(out[i])*(1-split_context)):] + if gen_text[i] != '': + split = re.split('[.!?]', gen_text[i]) + text = text.partition(list(filter(None, split))[-1])[-1] + + if truncate: + truncate_esc = re.escape(truncate) + if prefix and not include_prefix: + prefix_esc = re.escape(prefix) + pattern = '(?:{})(.*?)(?:{})'.format(prefix_esc, + truncate_esc) + else: + pattern = '(.*?)(?:{})'.format(truncate_esc) + + trunc_text = re.search(pattern, text, re.S) + if trunc_text: + text = trunc_text.group(1) + + if not truncated[i]: + gen_text[i] += text.lstrip('\n') + if trunc_text or (length is not None and total_tokens >= length-1): + truncated[i] = True for gen in gen_text: if destination_path: @@ -698,7 +698,7 @@ def cmd(): '--sample_delim', help="[generate] Delimiter between each generated sample.", nargs='?', default='=' * 20 + '\n', type=str) parser.add_argument( - '--split_context', help="[generate] When generating a sample of non-fixed length, feed this proportion of previous generation as context.", + '--split_context', help="[generate] When generating a sample longer than 1023 tokens, feed this proportion of previous generation as context.", nargs='?', default=0.5, type=float) # Positional arguments