Files
2026-07-06 17:41:08 -06:00

238 lines
9.6 KiB
Python

import os
import sys
import json
# ==============================================================================
# AUTO-CONFIG: Find and link hidden XLA compiler libraries inside the venv
# ==============================================================================
def configure_xla_paths():
venv_base = sys.prefix
for root, _, files in os.walk(venv_base):
if "libdevice.10.bc" in files:
cuda_dir = root.split("/nvvm")[0]
os.environ["XLA_FLAGS"] = f"--xla_gpu_cuda_data_dir={cuda_dir}"
print(f"[XLA Config] Successfully linked compiler data to: {cuda_dir}")
return True
print("[XLA Config] Warning: libdevice.10.bc not found. Training may crash.")
return False
configure_xla_paths()
import tensorflow as tf
#tf.debugging.set_log_device_placement(True)
import numpy as np
# ==========================================
# 1. ARCHITECTURE CONFIGURATION
# ==========================================
EMBEDDING_DIM = 512
RNN_UNITS = 512
# Pipeline & Training Hyperparameters
#TEXT_FILE = "combined_training_data.txt"
TEXT_FILE = "emily_post.txt"
VOCAB_FILE = "vocab.txt"
# DYNAMICALLY MANGLED DIRECTORY:
# This will automatically evaluate to something like: ./checkpoints_emb256_rnn512
CHECKPOINT_DIR = f"./checkpoints_emb{EMBEDDING_DIM}_rnn{RNN_UNITS}"
CONFIG_FILE = os.path.join(CHECKPOINT_DIR, "config.json")
SEQ_LENGTH = 512
BATCH_SIZE = 128
EPOCHS = 210
BUFFER_SIZE = 10000
SEED_TEXT = "In the case of unlawful enterance"
print("Num GPUs Available: ", len(tf.config.list_physical_devices('GPU')))
os.makedirs(CHECKPOINT_DIR, exist_ok=True)
# ==========================================
# 2. CONFIG CHECKPOINT GUARD (THE MAGIC SOUP)
# ==========================================
# If an existing config file is found, OVERWRITE your local variables
# to force the model to build at the size matching the saved weights.
if os.path.exists(CONFIG_FILE):
print(f"Found existing model configuration file at '{CONFIG_FILE}'.")
with open(CONFIG_FILE, 'r') as f:
saved_config = json.load(f)
EMBEDDING_DIM = saved_config["embedding_dim"]
RNN_UNITS = saved_config["rnn_units"]
print(f"-> Overrode network sizes to match saved profile: Embedding={EMBEDDING_DIM}, RNN Units={RNN_UNITS}")
else:
# Save the current configuration since it's a brand new run
current_config = {
"embedding_dim": EMBEDDING_DIM,
"rnn_units": RNN_UNITS
}
with open(CONFIG_FILE, 'w') as f:
json.dump(current_config, f, indent=4)
print(f"Created a new model configuration footprint file at '{CONFIG_FILE}'.")
# ==========================================
# 3. DATA PREPROCESSING
# ==========================================
if os.path.exists(VOCAB_FILE):
print(f"Loading existing vocabulary file from '{VOCAB_FILE}'...")
with open(VOCAB_FILE, 'r', encoding='utf-8') as f:
vocab = json.load(f)
else:
if not os.path.exists(TEXT_FILE):
print(f"File '{TEXT_FILE}' not found locally. Please run the process.py script!")
quit(1)
print(f"Loading raw text from {TEXT_FILE} to build vocabulary...")
with open(TEXT_FILE, 'r', encoding='utf-8') as f:
text = f.read()
temp_vectorizer = tf.keras.layers.TextVectorization(split="character", standardize="lower")
temp_vectorizer.adapt(tf.data.Dataset.from_tensor_slices([text]))
vocab = temp_vectorizer.get_vocabulary()
with open(VOCAB_FILE, 'w', encoding='utf-8') as f:
json.dump(vocab, f, ensure_ascii=False)
vocab_size = len(vocab)
vectorize_layer = tf.keras.layers.TextVectorization(
split="character", standardize="lower", vocabulary=vocab, output_mode="int"
)
dataset = None
if os.path.exists(TEXT_FILE):
with open(TEXT_FILE, 'r', encoding='utf-8') as f:
text = f.read()
all_ids = vectorize_layer(tf.constant([text]))[0]
ids_dataset = tf.data.Dataset.from_tensor_slices(all_ids)
sequences = ids_dataset.batch(SEQ_LENGTH + 1, drop_remainder=True)
def split_input_target(sequence):
return sequence[:-1], sequence[1:]
dataset = sequences.map(split_input_target).shuffle(BUFFER_SIZE).batch(BATCH_SIZE, drop_remainder=True).prefetch(tf.data.AUTOTUNE)
# ==========================================
# 4. MODEL ARCHITECTURE
# ==========================================
class CharacterTextModel(tf.keras.Model):
def __init__(self, vocab_size, embedding_dim, rnn_units, num_layers=1, dropout_rate=0.3):
super().__init__()
self.embedding = tf.keras.layers.Embedding(vocab_size, embedding_dim)
# Explicitly use return_state=True for ALL LSTM layers
self.lstm_layers = [
tf.keras.layers.LSTM(rnn_units, return_sequences=True, return_state=True)
for _ in range(num_layers)
]
self.dropout_layers = [tf.keras.layers.Dropout(dropout_rate) for _ in range(num_layers)]
self.dense = tf.keras.layers.Dense(vocab_size)
def call(self, inputs, states=None, return_state=False, training=False):
x = self.embedding(inputs, training=training)
# If no states are provided, initialize a list of None
if states is None:
states = [None] * len(self.lstm_layers)
new_states = []
for i, (lstm, dropout) in enumerate(zip(self.lstm_layers, self.dropout_layers)):
# LSTM returns: output, hidden state (h), and cell state (c)
x, h, c = lstm(x, initial_state=states[i], training=training)
x = dropout(x, training=training)
# Store both states as a list/tuple for this layer
new_states.append([h, c])
x = self.dense(x, training=training)
if return_state:
return x, new_states
return x
# Automatically scales to whatever dimensions were chosen or loaded!
model = CharacterTextModel(vocab_size=vocab_size, embedding_dim=EMBEDDING_DIM, rnn_units=RNN_UNITS)
model.build(input_shape=(BATCH_SIZE, SEQ_LENGTH))
import glob
# 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 = min(checkpoint_files, key=os.path.getmtime)
print(f"Restoring model layers from: {latest_checkpoint}")
model.load_weights(latest_checkpoint)
else:
print("Starting a clean initialization.")
loss = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True)
model.compile(optimizer='adam', loss=loss)
# ==========================================
# 5. SAMPLING & GENERATION LOGIC
# ==========================================
def produce_sample(model, seed, num_generate=300, temperature=0.7):
# Vectorize the initial seed text
input_chars = vectorize_layer(tf.constant([seed]))
input_ids = input_chars[0][:len(seed)].numpy().tolist()
generated_ids = []
states = None
# Step 1: "Warm up" the model with the seed text to build the initial states
# We pass the entire seed here
predictions, states = model(tf.expand_dims(input_ids, 0), states=states, return_state=True, training=False)
# Get the very last prediction from the seed sequence
predictions = predictions[0, -1, :] / temperature
predicted_id = tf.random.categorical(tf.expand_dims(predictions, 0), num_samples=1)[0, 0].numpy()
generated_ids.append(predicted_id)
# Step 2: Generation loop using ONLY the single newest token and updating states
for _ in range(num_generate - 1):
# Pass ONLY the last predicted token, plus the existing states
predictions, states = model(tf.expand_dims([predicted_id], 0), states=states, return_state=True, training=False)
# Scale by temperature and sample
predictions = predictions[0, -1, :] / temperature
predicted_id = tf.random.categorical(tf.expand_dims(predictions, 0), num_samples=1)[0, 0].numpy()
generated_ids.append(predicted_id)
return "".join([vocab[idx] for idx in generated_ids])
class GenerationCallback(tf.keras.callbacks.Callback):
def on_epoch_end(self, epoch, logs=None):
print(f"\n--- End of Epoch {epoch+1} ---")
if epoch % 10 == 0:
for temp in [0.3, 0.6, 0.9, 1.2]:
print(f"Seed: \"{SEED_TEXT}\" -> {produce_sample(self.model, seed=SEED_TEXT, num_generate=200, temperature=temp)}\n")
checkpoint_prefix = os.path.join(CHECKPOINT_DIR, "ckpt_{epoch}.weights.h5")
checkpoint_callback = tf.keras.callbacks.ModelCheckpoint(filepath=checkpoint_prefix, save_weights_only=True)
# Stop training when validation loss stops improving for 3 epochs
early_stopping = tf.keras.callbacks.EarlyStopping(
monitor='loss',
patience=3,
restore_best_weights=True # Automatically rolls back best weights
)
# Create the learning rate scheduler callback
lr_scheduler = tf.keras.callbacks.ReduceLROnPlateau(
monitor='loss', # Can change to 'val_loss' if using a validation split
factor=0.5, # Multiply the learning rate by 0.5 when triggered (cuts it in half)
patience=2, # Number of epochs with no improvement before dropping LR
min_lr=1e-6, # Don't let the learning rate drop lower than this
verbose=1 # Prints a message when the learning rate changes
)
# ==========================================
# 6. START RUN
# ==========================================
if dataset is not None:
model.fit(dataset, epochs=EPOCHS, callbacks=[GenerationCallback(), checkpoint_callback, early_stopping, lr_scheduler])
else:
print(f"\n--- Inference Mode ({EMBEDDING_DIM}dim, {RNN_UNITS}units) ---")
print(f"Result: {produce_sample(model, seed=SEED_TEXT, num_generate=300, temperature=0.6)}")