Commit 4ef0a794 authored by Rudolf Chrispens's avatar Rudolf Chrispens
Browse files

refactoring the lint stuff understanding the batch training

parent bb351124
Loading
Loading
Loading
Loading
+1 −1
Original line number Diff line number Diff line
@@ -9,6 +9,6 @@
    ],
    "editor.minimap.enabled": false,
    "python.linting.flake8Args": [
        "--ignore=E302,E303,E304,E305,E501",
        "--ignore=E302,E303,E304,E305,E501,E251,E128",
    ],
}
 No newline at end of file
+23 −10
Original line number Diff line number Diff line
# !/usr/bin/env python3

import tensorflow as tf
import numpy as np
from lstm.configuration import current as config
@@ -31,8 +30,8 @@ def train():
    decoder_targets = tf.placeholder(shape=(None, None), dtype=tf.int32, name='decoder_targets')

    # embeddings
    embeddings = tf.Variable(tf.random_uniform([config.trainer.vocab_size, config.trainer.input_embedding_size], -1.0, 1.0), dtype=tf.float32)
    encoder_inputs_embedded = tf.nn.embedding_lookup(embeddings, encoder_inputs)
    embeddings = tf.Variable(tf.random_uniform([config.trainer.vocab_size, config.trainer.input_embedding_size], -1.0, 1.0), dtype=tf.float32, name='embeddings')
    encoder_inputs_embedded = tf.nn.embedding_lookup(embeddings, encoder_inputs)  # TODO alternative embedding_ops.embedding_lookup(embeddings, encoder_inputs)

    # define encoder
    encoder_cell = LSTMCell(config.trainer.encoder_hidden_units)  # each cell each neuron is an lstm itself!
@@ -41,11 +40,13 @@ def train():
    # normal rnn only takes the past into account a dynamic RNN does take the future into account
    # biodirectional LSTM
    ((encoder_fw_outputs, encoder_bw_outputs), (encoder_fw_final_state, encoder_bw_final_state)) = (
        tf.nn.bidirectional_dynamic_rnn(cell_fw=encoder_cell,
        tf.nn.bidirectional_dynamic_rnn(
            cell_fw = encoder_cell,
            cell_bw = encoder_cell,
            inputs = encoder_inputs_embedded,
            sequence_length = encoder_inputs_length,
        dtype=tf.float32, time_major=True)
            dtype = tf.float32,
            time_major = True)
    )

    # bidirectional step1
@@ -68,7 +69,7 @@ def train():
    # output projects
    # weights and biases
    # SOFT ATTENTION
    W = tf.Variable(tf.random_uniform([config.trainer.decoder_hidden_units], vocab_size), -1, 1), dtype=tf.float32)
    W = tf.Variable(tf.random_uniform([config.trainer.decoder_hidden_units], vocab_size), -1, 1, dtype = tf.float32)
    b = tf.Variable(tf.zeroes([config.trainer.vocab_size]), dtype = tf.float32)


@@ -103,7 +104,7 @@ def train():
    # cross entropy loss
    # one hot encode the target values so we dont rank just differentiate
    stepwise_cross_entropy = tf.nn.softmax_cross_entropy_with_logits(
        labels=tf.one_hot)decoder_targets, depth=config.trainer.vocab_size, dtype=tf.float32),
        labels=tf.one_hot(decoder_targets, depth=config.trainer.vocab_size, dtype=tf.float32),
        logits=decoder_logits
    )

@@ -117,8 +118,8 @@ def train():
    # training the real deal executes here:
    try:
        for batch in range(config.trainer.max_batches):
            fd = # TODO get a sequence to learn from
            _, l = sess.run([train_op], loss), fd)
            fd =  # TODO get a sequence to learn from seq2seq model next_feed()
            _, l = sess.run((([train_op], loss), fd)
            loss_track.append(l)

            if(batch == 0 or batch % config.trainer.batches_in_epoch == 0)
@@ -132,7 +133,6 @@ def train():
                    if(i >= 2):
                        break
                print()

    except KeyboardInterrupt:
        print('\n.\n.\n.\n...training interupted')

@@ -140,6 +140,19 @@ def train():
# manually spcifying loop function over time -  to get initial cell state and input to RNN
# normally we would just use dynamic_rnn, but lets get detailed here with raw_rnn


def next_feed():
    batch = next(batches)
    encoder_inputs_, encoder_input_lengths_ = helpers.batch(batch)
    decoder_targets_, _ = helpers.batch(
        [(sequence) + [EOS] + [PAD] * 2 for sequence in batch]
    )
    return {
        encoder_inputs: encoder_inputs_,
        encoder_inputs_length: encoder_input_lengths_,
        decoder_targets: decoder_targets_,
    }

# we define and return these values, no operations occure here
def loop_fn_initial():
    initial_elements_finished = (0 >= decoder_lengths)