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 import numpy as np # ========================================== # 1. ARCHITECTURE CONFIGURATION # ========================================== EMBEDDING_DIM = 256 RNN_UNITS = 512 # Pipeline & Training Hyperparameters TEXT_FILE = "combined_training_data.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 = 150 BATCH_SIZE = 64 EPOCHS = 100 BUFFER_SIZE = 10000 SEED_TEXT = "it was a dark and stormy night" 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): # super().__init__() # self.embedding = tf.keras.layers.Embedding(vocab_size, embedding_dim) # self.lstm1 = tf.keras.layers.LSTM(rnn_units, return_sequences=True, return_state=True) # self.lstm2 = tf.keras.layers.LSTM(rnn_units, return_sequences=True, return_state=True) # 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 states is None: # state_1, state_2 = None, None # else: # state_1, state_2 = states # x, h1, c1 = self.lstm1(x, initial_state=state_1, training=training) # x, h2, c2 = self.lstm2(x, initial_state=state_2, training=training) # x = self.dense(x, training=training) # if return_state: # return x, [(h1, c1), (h2, c2)] # return x class CharacterTextModel(tf.keras.Model): def __init__(self, vocab_size, embedding_dim, rnn_units, dropout_rate=0.2): super().__init__() self.embedding = tf.keras.layers.Embedding(vocab_size, embedding_dim) # Layer 1 self.lstm1 = tf.keras.layers.LSTM(rnn_units, return_sequences=True, return_state=True) self.ln1 = tf.keras.layers.LayerNormalization() self.dropout1 = tf.keras.layers.Dropout(dropout_rate) # Layer 2 self.lstm2 = tf.keras.layers.LSTM(rnn_units, return_sequences=True, return_state=True) self.ln2 = tf.keras.layers.LayerNormalization() self.dropout2 = tf.keras.layers.Dropout(dropout_rate) # Layer 3 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.dropout_final = tf.keras.layers.Dropout(dropout_rate) # Final output dense layer 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) # Unpack states cleanly for 3 layers if states is None: state_1, state_2, state_3, state_4, state_5, state_6, state_final = None, None, None, None, None, None, None else: state_1, state_2, state_3, state_4, state_5, state_6, state_final = states # Pass through Layer 1 x, h1, c1 = self.lstm1(x, initial_state=state_1, 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) # Output logits x = self.dense(x, training=training) if return_state: return x, [(h1, c1), (h2, c2), (h3, c3), (h4, c4), (h5, c5), (h6, c6), (h_final, c_final)] 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)) latest_checkpoint = tf.train.latest_checkpoint(CHECKPOINT_DIR) if latest_checkpoint: 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): input_chars = vectorize_layer(tf.constant([seed])) input_ids = input_chars[0][:len(seed)].numpy().tolist() generated_ids, states = [], None for _ in range(num_generate): current_tokens = input_ids[-SEQ_LENGTH:] predictions, states = model(tf.expand_dims(current_tokens, 0), states=states, return_state=True, training=False) 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) input_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.1, 0.4, 0.7]: 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) # ========================================== # 6. START RUN # ========================================== if dataset is not None: model.fit(dataset, epochs=EPOCHS, callbacks=[GenerationCallback(), checkpoint_callback]) 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)}")