diff --git a/gpt_2_simple/gpt_2.py b/gpt_2_simple/gpt_2.py index e02021b..96699c1 100644 --- a/gpt_2_simple/gpt_2.py +++ b/gpt_2_simple/gpt_2.py @@ -362,7 +362,8 @@ 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 +379,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' @@ -387,57 +392,76 @@ 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 = tf.placeholder(tf.int32, [batch_size, None]) - context_tokens = enc.encode(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) - 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: + gen_text = [''] * batch_size + 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: batch_size * [context_tokens] + context: 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 - 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, gen_text, re.S) - if trunc_text: - gen_text = trunc_text.group(1) - gen_text = gen_text.lstrip('\n') + + total_tokens += num_tokens + + for i in range(batch_size): + text = enc.decode(out[i]) + if prefix: + 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: - 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 longer than 1023 tokens, 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 )