rebasing to main
This commit is contained in:
@@ -19,12 +19,13 @@ def configure_xla_paths():
|
|||||||
configure_xla_paths()
|
configure_xla_paths()
|
||||||
|
|
||||||
import tensorflow as tf
|
import tensorflow as tf
|
||||||
|
#tf.debugging.set_log_device_placement(True)
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
# ==========================================
|
# ==========================================
|
||||||
# 1. ARCHITECTURE CONFIGURATION
|
# 1. ARCHITECTURE CONFIGURATION
|
||||||
# ==========================================
|
# ==========================================
|
||||||
EMBEDDING_DIM = 256
|
EMBEDDING_DIM = 1024
|
||||||
RNN_UNITS = 512
|
RNN_UNITS = 512
|
||||||
|
|
||||||
# Pipeline & Training Hyperparameters
|
# Pipeline & Training Hyperparameters
|
||||||
@@ -36,8 +37,8 @@ VOCAB_FILE = "vocab.txt"
|
|||||||
CHECKPOINT_DIR = f"./checkpoints_emb{EMBEDDING_DIM}_rnn{RNN_UNITS}"
|
CHECKPOINT_DIR = f"./checkpoints_emb{EMBEDDING_DIM}_rnn{RNN_UNITS}"
|
||||||
CONFIG_FILE = os.path.join(CHECKPOINT_DIR, "config.json")
|
CONFIG_FILE = os.path.join(CHECKPOINT_DIR, "config.json")
|
||||||
|
|
||||||
SEQ_LENGTH = 150
|
SEQ_LENGTH = 300
|
||||||
BATCH_SIZE = 64
|
BATCH_SIZE = 128
|
||||||
EPOCHS = 100
|
EPOCHS = 100
|
||||||
BUFFER_SIZE = 10000
|
BUFFER_SIZE = 10000
|
||||||
SEED_TEXT = "it was a dark and stormy night"
|
SEED_TEXT = "it was a dark and stormy night"
|
||||||
@@ -132,108 +133,63 @@ if os.path.exists(TEXT_FILE):
|
|||||||
# return x
|
# return x
|
||||||
|
|
||||||
class CharacterTextModel(tf.keras.Model):
|
class CharacterTextModel(tf.keras.Model):
|
||||||
def __init__(self, vocab_size, embedding_dim, rnn_units, dropout_rate=0.2):
|
def __init__(self, vocab_size, embedding_dim, rnn_units, num_layers=2, dropout_rate=0.2):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.embedding = tf.keras.layers.Embedding(vocab_size, embedding_dim)
|
self.embedding = tf.keras.layers.Embedding(vocab_size, embedding_dim)
|
||||||
|
|
||||||
# Layer 1
|
# Create paired lists of LSTMs and Dropouts based on your desired depth
|
||||||
self.lstm1 = tf.keras.layers.LSTM(rnn_units, return_sequences=True, return_state=True)
|
self.lstm_layers = []
|
||||||
self.ln1 = tf.keras.layers.LayerNormalization()
|
self.dropout_layers = []
|
||||||
self.dropout1 = tf.keras.layers.Dropout(dropout_rate)
|
|
||||||
|
|
||||||
# Layer 2
|
# If num_layers=3, this loop runs 2 times, leaving the 3rd layer as the final layer
|
||||||
self.lstm2 = tf.keras.layers.LSTM(rnn_units, return_sequences=True, return_state=True)
|
for _ in range(num_layers - 1):
|
||||||
self.ln2 = tf.keras.layers.LayerNormalization()
|
self.lstm_layers.append(
|
||||||
self.dropout2 = tf.keras.layers.Dropout(dropout_rate)
|
tf.keras.layers.LSTM(rnn_units, return_sequences=True)
|
||||||
|
)
|
||||||
|
self.dropout_layers.append(
|
||||||
|
tf.keras.layers.Dropout(dropout_rate)
|
||||||
|
)
|
||||||
|
|
||||||
# Layer 3
|
# Keep return_state ONLY on the final layer if needed for text generation
|
||||||
self.lstm3 = tf.keras.layers.LSTM(rnn_units, return_sequences=True, return_state=True)
|
|
||||||
self.ln3 = tf.keras.layers.LayerNormalization()
|
|
||||||
self.dropout3 = tf.keras.layers.Dropout(dropout_rate)
|
|
||||||
|
|
||||||
# Layer 4
|
|
||||||
self.lstm4 = tf.keras.layers.LSTM(rnn_units, return_sequences=True, return_state=True)
|
|
||||||
self.ln4 = tf.keras.layers.LayerNormalization()
|
|
||||||
self.dropout4 = tf.keras.layers.Dropout(dropout_rate)
|
|
||||||
|
|
||||||
# Layer 5
|
|
||||||
self.lstm5 = tf.keras.layers.LSTM(rnn_units, return_sequences=True, return_state=True)
|
|
||||||
self.ln5 = tf.keras.layers.LayerNormalization()
|
|
||||||
self.dropout5 = tf.keras.layers.Dropout(dropout_rate)
|
|
||||||
|
|
||||||
# Layer 6
|
|
||||||
self.lstm6 = tf.keras.layers.LSTM(rnn_units, return_sequences=True, return_state=True)
|
|
||||||
self.ln6 = tf.keras.layers.LayerNormalization()
|
|
||||||
self.dropout6 = tf.keras.layers.Dropout(dropout_rate)
|
|
||||||
|
|
||||||
# Layer _final (Added for deeper text/context comprehension)
|
|
||||||
self.lstm_final = tf.keras.layers.LSTM(rnn_units, return_sequences=True, return_state=True)
|
self.lstm_final = tf.keras.layers.LSTM(rnn_units, return_sequences=True, return_state=True)
|
||||||
self.dropout_final = tf.keras.layers.Dropout(dropout_rate)
|
self.dropout_final = tf.keras.layers.Dropout(dropout_rate)
|
||||||
|
|
||||||
# Final output dense layer
|
|
||||||
self.dense = tf.keras.layers.Dense(vocab_size)
|
self.dense = tf.keras.layers.Dense(vocab_size)
|
||||||
|
|
||||||
def call(self, inputs, states=None, return_state=False, training=False):
|
def call(self, inputs, states=None, return_state=False, training=False):
|
||||||
x = self.embedding(inputs, training=training)
|
x = self.embedding(inputs, training=training)
|
||||||
|
|
||||||
# Unpack states cleanly for 3 layers
|
# Loop through both lists simultaneously using zip()
|
||||||
if states is None:
|
for lstm, dropout in zip(self.lstm_layers, self.dropout_layers):
|
||||||
state_1, state_2, state_3, state_4, state_5, state_6, state_final = None, None, None, None, None, None, None
|
x = lstm(x, training=training)
|
||||||
else:
|
x = dropout(x, training=training)
|
||||||
state_1, state_2, state_3, state_4, state_5, state_6, state_final = states
|
|
||||||
|
|
||||||
# Pass through Layer 1
|
# Final LSTM Layer
|
||||||
x, h1, c1 = self.lstm1(x, initial_state=state_1, training=training)
|
x, h, c = self.lstm_final(x, initial_state=states, training=training)
|
||||||
x = self.ln1(x, training=training)
|
|
||||||
x = self.dropout1(x, training=training)
|
|
||||||
|
|
||||||
# Pass through Layer 2
|
|
||||||
x, h2, c2 = self.lstm2(x, initial_state=state_2, training=training)
|
|
||||||
x = self.ln2(x, training=training)
|
|
||||||
x = self.dropout2(x, training=training)
|
|
||||||
|
|
||||||
# Pass through Layer 3
|
|
||||||
x, h3, c3 = self.lstm3(x, initial_state=state_3, training=training)
|
|
||||||
x = self.ln3(x, training=training)
|
|
||||||
x = self.dropout3(x, training=training)
|
|
||||||
|
|
||||||
# Pass through Layer 4
|
|
||||||
x, h4, c4 = self.lstm4(x, initial_state=state_4, training=training)
|
|
||||||
x = self.ln4(x, training=training)
|
|
||||||
x = self.dropout4(x, training=training)
|
|
||||||
|
|
||||||
# Pass through Layer 5
|
|
||||||
x, h5, c5 = self.lstm5(x, initial_state=state_5, training=training)
|
|
||||||
x = self.ln5(x, training=training)
|
|
||||||
x = self.dropout5(x, training=training)
|
|
||||||
|
|
||||||
# Pass through Layer 6
|
|
||||||
x, h6, c6 = self.lstm6(x, initial_state=state_6, training=training)
|
|
||||||
x = self.ln6(x, training=training)
|
|
||||||
x = self.dropout6(x, training=training)
|
|
||||||
|
|
||||||
# Pass through Layer _final
|
|
||||||
x, h_final, c_final = self.lstm_final(x, initial_state=state_final, training=training)
|
|
||||||
x = self.dropout_final(x, training=training)
|
x = self.dropout_final(x, training=training)
|
||||||
|
|
||||||
# Output logits
|
|
||||||
x = self.dense(x, training=training)
|
x = self.dense(x, training=training)
|
||||||
|
|
||||||
if return_state:
|
if return_state:
|
||||||
return x, [(h1, c1), (h2, c2), (h3, c3), (h4, c4), (h5, c5), (h6, c6), (h_final, c_final)]
|
return x, (h, c)
|
||||||
return x
|
return x
|
||||||
|
|
||||||
# Automatically scales to whatever dimensions were chosen or loaded!
|
# Automatically scales to whatever dimensions were chosen or loaded!
|
||||||
model = CharacterTextModel(vocab_size=vocab_size, embedding_dim=EMBEDDING_DIM, rnn_units=RNN_UNITS)
|
model = CharacterTextModel(vocab_size=vocab_size, embedding_dim=EMBEDDING_DIM, rnn_units=RNN_UNITS)
|
||||||
model.build(input_shape=(BATCH_SIZE, SEQ_LENGTH))
|
model.build(input_shape=(BATCH_SIZE, SEQ_LENGTH))
|
||||||
|
|
||||||
latest_checkpoint = tf.train.latest_checkpoint(CHECKPOINT_DIR)
|
import glob
|
||||||
if latest_checkpoint:
|
|
||||||
|
# Find all files matching the pattern
|
||||||
|
checkpoint_files = glob.glob(os.path.join(CHECKPOINT_DIR, "ckpt_*.weights.h5"))
|
||||||
|
|
||||||
|
if checkpoint_files:
|
||||||
|
# Sort files naturally or by modification time to get the latest one
|
||||||
|
latest_checkpoint = max(checkpoint_files, key=os.path.getmtime)
|
||||||
print(f"Restoring model layers from: {latest_checkpoint}")
|
print(f"Restoring model layers from: {latest_checkpoint}")
|
||||||
model.load_weights(latest_checkpoint)
|
model.load_weights(latest_checkpoint)
|
||||||
else:
|
else:
|
||||||
print("Starting a clean initialization.")
|
print("Starting a clean initialization.")
|
||||||
|
|
||||||
|
|
||||||
loss = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True)
|
loss = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True)
|
||||||
model.compile(optimizer='adam', loss=loss)
|
model.compile(optimizer='adam', loss=loss)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user