Aoban1.1A-Refined / Python /ai thingy.py
Aobangaming's picture
Upload ai thingy.py
9a2aa07 verified
Raw History Blame
28.2 kB
import torch
import torch.nn as nn
from torch.optim import Adam
from torch.utils.data import Dataset, DataLoader
import math
import time
import csv
import os
import re
# --- HYPERPARAMETERS ---
D_MODEL = 256
NUM_HEADS = 4
NUM_LAYERS = 6
DROPOUT = 0.1
MAX_SEQ_LENGTH = 32
LEARNING_RATE = 1e-4
NUM_EPOCHS = 20 # Default full training epochs
BATCH_SIZE = 32
# Increase training epochs as requested
NUM_EPOCHS = 25 # Default full training epochs (increased)
INTERACTIVE_EPOCHS = 15 # Epochs for quick retraining (increased)
# --- GENERATION SETTINGS ---
TOP_K = 3
REPETITION_PENALTY = 3
TEMPERATURE = 1
# --- PERSISTENCE CONFIGURATION ---
DATA_FILE = 'training_data.csv' # File where all training data is stored
# --- INITIAL DATA FALLBACK (The 27 sentences you provided) ---
DEFAULT_TRAINING_DATA = [
"The quick brown fox jumps over the lazy dog.",
"A glass of water is clear.",
"The sun is shining bright and the sky is clear.",
"The dog and the fox are friends forever.",
"Coding with Pytorch and Transformers is fun and very rewarding.",
"A computer runs very fast and never stops.",
"The windows are big and bright.",
"A green park is a great place to relax.",
"The sky is clear today, with no clouds.",
"The cat jumped over the fence.",
"The plane has many windows.",
"A big bird flew over the house.",
"The plane smoothly landed on the concrete runway.",
"The bird flew above the bustling city.",
"The plane had an engine failure and had to land in the river.",
"The Cessna 172 is a low-wing monoplane.",
"The plane flew by the trees.",
"The plane, almost out of fuel, finally landed at an airport.",
"The angry bird flew away furiously.",
"A plane is a machine that flies.",
"The fast plane landed at the bright airport.",
"The plane quickly landed on the runway.",
"The letter A is part of the alphabet.",
"The plane landed hardly on a grass runway in the forest.",
"The clouds were floating above the ground.",
"The plane was a very bright plane, it's livery glimmered in the night sky.",
"The GPWS sounds on a plane are like Caution Terrain PULL up PULL up."
]
# --- FILE I/O FUNCTIONS (CRITICAL FOR PERSISTENCE) ---
def load_data_from_csv(filepath):
"""Loads all training sentences from the CSV file, or returns the default data."""
texts = []
def split_into_sentences(paragraph):
# Split on sentence end punctuation followed by whitespace and a capital or number
# Use a safe regex string and fall back to newline/sentence punctuation splitting on error
try:
pattern = r'(?<=[\.\!?])\s+(?=[A-Z0-9"\'""\u201c])'
parts = re.split(pattern, paragraph)
return [p.strip() for p in parts if p and p.strip()]
except re.error:
# fallback: split on sentence enders and newlines
parts = re.split(r'[\.\!?]\s+|\n+', paragraph)
return [p.strip() for p in parts if p and p.strip()]
# Attempt to read existing data
if os.path.exists(filepath) and os.path.getsize(filepath) > 0:
print(f"[SYSTEM] Loading training data from {filepath}...")
try:
with open(filepath, 'r', newline='', encoding='utf-8') as f:
reader = csv.reader(f)
raw_rows = []
for row in reader:
if row and row[0].strip():
raw_text = row[0].strip()
# Remove surrounding quotes if present
if (raw_text.startswith('"') and raw_text.endswith('"')) or (raw_text.startswith("'") and raw_text.endswith("'")):
raw_text = raw_text[1:-1].strip()
if raw_text:
raw_rows.append(raw_text)
# Now split rows into sentences, filter and handle adjacent runs
sequence = []
for raw in raw_rows:
# If the row contains multiple sentences, split them
parts = split_into_sentences(raw)
# If splitting produced only one part but it contains multiple internal newlines, also split on newlines
if len(parts) == 1 and '\n' in parts[0]:
parts = [p.strip() for p in parts[0].splitlines() if p.strip()]
for s in parts:
# Normalize whitespace and strip quotes
s_clean = ' '.join(s.split()).strip(' "\'')
words = s_clean.split()
# Basic length filters to remove garbage/too-short sentences
if len(words) < 5:
continue
if len(words) > 200:
# skip extremely long paragraphs
continue
# Filter out noisy/corrupted lines
# Skip if contains excessive repetition (same word 3+ times in a row)
is_noisy = False
for i in range(len(words) - 2):
if words[i] == words[i+1] == words[i+2]:
is_noisy = True
break
if is_noisy:
continue
# Skip lines that look like training artifacts (high ratio of common junk words)
junk_patterns = ['pull', 'up', 'land', 'river', 'sky', 'clear', 'table']
junk_count = sum(1 for w in words if w in junk_patterns)
if junk_count > len(words) * 0.4: # more than 40% junk
continue
sequence.append(s_clean)
# Collapse consecutive identical sentences (runs) to at most two copies
i = 0
while i < len(sequence):
j = i + 1
while j < len(sequence) and sequence[j] == sequence[i]:
j += 1
run_len = j - i
if run_len == 1:
texts.append(sequence[i])
else:
# keep first and last occurrence of the run
texts.append(sequence[i])
texts.append(sequence[i])
i = j
except Exception as e:
print(f"[ERROR] Error loading CSV: {e}. Falling back to default data.")
texts = [] # Clear corrupted load
# If no data loaded (file missing, empty, or corrupted), use the default knowledge base
if not texts:
print("[SYSTEM] CSV file not found or empty. Using default knowledge base.")
return list(DEFAULT_TRAINING_DATA)
# Debug: report how many sentences were actually loaded and sample content
print(f"[SYSTEM] Loaded {len(texts)} sentence(s) from {filepath}.")
sample_head = texts[:10]
sample_tail = texts[-10:]
print("[DEBUG] First loaded sentences:")
for i, s in enumerate(sample_head, 1):
print(f" {i}: {s[:200]}")
if len(texts) > 10:
print("[DEBUG] Last loaded sentences:")
start_index = max(0, len(texts) - 10)
for i, s in enumerate(texts[start_index:], start_index + 1):
print(f" {i}: {s[:200]}")
return texts
def save_data_to_csv(filepath, texts):
"""Saves the entire list of training sentences (including new ones) to the CSV."""
# The 'w' mode ensures the file is overwritten with the complete, updated dataset.
print(f"[SYSTEM] Saving {len(texts)} sentences to {filepath} using 'w' mode...")
try:
with open(filepath, 'w', newline='', encoding='utf-8') as f:
writer = csv.writer(f)
# Write each sentence as a single row/column entry
for text in texts:
writer.writerow([text])
except Exception as e:
print(f"[ERROR] Error saving to CSV: {e}")
# --- TOKENIZER ---
class SimpleTokenizer:
def __init__(self, texts):
self.word_to_idx = {"<PAD>": 0, "<UNK>": 1}
self.idx_to_word = {0: "<PAD>", 1: "<UNK>"}
self.build_vocab(texts)
def build_vocab(self, texts):
for text in texts:
for word in text.lower().split():
word = word.strip(".,!?")
if word not in self.word_to_idx:
idx = len(self.word_to_idx)
self.word_to_idx[word] = idx
self.idx_to_word[idx] = word
def encode(self, text, max_len):
words = [word.strip(".,!?") for word in text.lower().split()]
indices = [self.word_to_idx.get(word, self.word_to_idx["<UNK>"]) for word in words]
# Padding and Truncation
if len(indices) < max_len:
indices.extend([self.word_to_idx["<PAD>"]] * (max_len - len(indices)))
elif len(indices) > max_len:
indices = indices[:max_len]
return torch.tensor(indices, dtype=torch.long)
def decode(self, indices):
return " ".join([self.idx_to_word.get(idx.item(), "<UNK>") for idx in indices if idx.item() != self.word_to_idx["<PAD>"]])
@property
def vocab_size(self):
return len(self.word_to_idx)
# --- DATASET ---
class TextDataset(Dataset):
def __init__(self, texts, tokenizer, max_len):
self.data = []
for text in texts:
encoded = tokenizer.encode(text, max_len)
self.data.append(encoded)
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return self.data[idx]
# --- TRANSFORMER MODEL COMPONENTS (UNMODIFIED) ---
class PositionalEncoding(nn.Module):
def __init__(self, d_model, max_len=5000):
super(PositionalEncoding, self).__init__()
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
pe = pe.unsqueeze(0).transpose(0, 1)
self.register_buffer('pe', pe)
def forward(self, x):
return x + self.pe[:x.size(1), :].transpose(0, 1)
class TransformerLanguageModel(nn.Module):
def __init__(self, vocab_size, d_model, nhead, num_layers, dropout, max_len):
super(TransformerLanguageModel, self).__init__()
self.model_type = 'Transformer'
self.d_model = d_model
self.vocab_size = vocab_size
self.embedding = nn.Embedding(vocab_size, d_model)
self.pos_encoder = PositionalEncoding(d_model, max_len)
# Use decoder layers for proper causal masking in text generation
decoder_layer = nn.TransformerDecoderLayer(
d_model=d_model,
nhead=nhead,
dim_feedforward=d_model*4,
dropout=dropout,
batch_first=True
)
self.transformer_decoder = nn.TransformerDecoder(decoder_layer, num_layers=num_layers)
self.fc_out = nn.Linear(d_model, vocab_size)
self.init_weights()
def init_weights(self):
initrange = 0.1
self.embedding.weight.data.uniform_(-initrange, initrange)
self.fc_out.bias.data.zero_()
self.fc_out.weight.data.uniform_(-initrange, initrange)
def forward(self, src):
src = self.embedding(src) * math.sqrt(self.d_model)
src = self.pos_encoder(src)
# Create causal mask to prevent attending to future tokens
seq_len = src.size(1)
causal_mask = torch.triu(torch.ones(seq_len, seq_len, device=src.device) * float('-inf'), diagonal=1)
# Decoder expects (tgt, memory) but we use same for both (causal language modeling)
output = self.transformer_decoder(src, src, tgt_mask=causal_mask)
return self.fc_out(output)
# --- TRAINING FUNCTIONS ---
def train_model(model, data_loader, optimizer, criterion, device, epochs):
model.train()
for epoch in range(1, epochs + 1):
total_loss = 0.0
start_time = time.time()
for batch in data_loader:
batch = batch.to(device)
src = batch[:, :-1]
tgt = batch[:, 1:]
optimizer.zero_grad()
output = model(src)
loss = criterion(output.reshape(-1, output.size(-1)), tgt.reshape(-1))
loss.backward()
optimizer.step()
total_loss += loss.item()
avg_loss = total_loss / len(data_loader)
# Print update frequency
if epochs > 10 and epoch % (epochs // 10) == 0:
print(f"Epoch {epoch}/{epochs}, Average Loss: {avg_loss:.4f}")
elif epochs > 0 and epochs <= 50 and epoch % 10 == 0:
print(f"Epoch {epoch}/{epochs}, Average Loss: {avg_loss:.4f}")
if epochs > 0:
print(f"[TRAINING COMPLETE] Model weights have been updated.")
# --- GENERATION FUNCTION (UNMODIFIED) ---
def generate_text(model, tokenizer, prompt, max_len, device, top_k=40, penalty=1.8, temperature=1.0):
model.eval()
encoded_prompt = tokenizer.encode(prompt, max_len=max_len).to(device)
# Count non-PAD tokens in encoded prompt to get true prompt length
pad_idx = tokenizer.word_to_idx["<PAD>"]
prompt_len = (encoded_prompt != pad_idx).sum().item()
generated_indices = encoded_prompt[:prompt_len].tolist()
input_ids = encoded_prompt.unsqueeze(0)
for i in range(prompt_len, max_len):
src_input = input_ids[:, :i]
with torch.no_grad():
output = model(src_input)
logits = output[0, i-1, :]
# --- START GRAMMATICAL FILTERING INJECTION ---
# ... (The rest of the code continues with: logits = logits / TEMPERATURE, TOP_K filtering, etc.)
# Apply Repetition Penalty
history = generated_indices
for idx in set(history):
if logits[idx] > 0:
logits[idx] /= penalty
else:
logits[idx] *= penalty
# Apply temperature scaling before top-k
logits = logits / temperature
# Apply Top-K Sampling
top_k_values, top_k_indices = torch.topk(logits, min(top_k, len(logits)))
probabilities = torch.softmax(top_k_values, dim=0)
try:
next_token_idx = torch.multinomial(probabilities, num_samples=1).item()
except RuntimeError:
predicted_token = top_k_indices[0].item()
if predicted_token == tokenizer.word_to_idx["<PAD>"]:
break
else:
predicted_token = top_k_indices[next_token_idx].item()
generated_indices.append(predicted_token)
input_ids[0, i] = predicted_token
# --- START USER-REQUESTED WERE PLURALIZATION RULE ---
# 1. Decode the word the model just selected
next_word = current_tokenizer.idx_to_word.get(predicted_token, '<UNK>').lower().strip(".,!?")
# 2. Check the condition: If the predicted word is 'were'
if next_word == 'were':
# Identify the singular noun/word that precedes 'were'
last_token_id = generated_indices[-1]
last_word_str = current_tokenizer.idx_to_word.get(last_token_id, '<UNK>').lower().strip(".,!?")
# Form the plural string (simple 's' rule)
plural_word_str = last_word_str + 's'
# Find the token ID for the plural word in the vocabulary
plural_token_id = current_tokenizer.word_to_idx.get(plural_word_str)
if plural_token_id is not None:
# If the plural form exists, replace the previous word's ID in the sentence
generated_indices[-1] = plural_token_id
print(f"[RULE ENFORCED] Modified '{last_word_str}' to '{plural_word_str}' before 'were' was added.")
# If the plural form is not in the vocabulary, we simply skip the modification.
# --- END USER-REQUESTED WERE PLURALIZATION RULE ---
# Decode only the continuation text
decoded_text = tokenizer.decode(torch.tensor(generated_indices, dtype=torch.long))
prompt_words = [word.strip(".,!?") for word in prompt.lower().split()]
decoded_words = decoded_text.split()
start_index = len(prompt_words)
continuation_text = " ".join(decoded_words[start_index:])
return continuation_text.replace(" <pad>", "").strip()
# --- MAIN EXECUTION ---
# Global variables for model/tokenizer instances
last_generated_text = None
last_user_prompt = None
current_tokenizer = None
current_model = None
device = torch.device("cpu")
live_data_updates = [] # Temporary queue for new sentences added during the current session
initial_training_texts = [] # Stores all data loaded from CSV
def initialize_or_retrain(initial_train=True, use_live_data=False, epochs=NUM_EPOCHS):
global current_tokenizer, current_model, live_data_updates, initial_training_texts
# 1. Load Data (Permanent)
if initial_train:
initial_training_texts = load_data_from_csv(DATA_FILE)
training_data = list(initial_training_texts)
# 2. Add Live Data
if use_live_data:
print(f"[SYSTEM] Retraining on {len(initial_training_texts)} base examples plus {len(live_data_updates)} new examples.")
training_data.extend(live_data_updates)
# 3. Tokenizer Initialization and Model Rebuild if necessary
old_vocab_size = current_tokenizer.vocab_size if current_tokenizer else 0
current_tokenizer = SimpleTokenizer(training_data)
new_vocab_size = current_tokenizer.vocab_size
if new_vocab_size != old_vocab_size or initial_train:
if initial_train:
print(f"Tokenizer Vocabulary Size: {new_vocab_size}")
print(f"\nModel D_MODEL={D_MODEL}, NUM_HEADS={NUM_HEADS}, NUM_LAYERS={NUM_LAYERS}")
current_model = TransformerLanguageModel(
vocab_size=new_vocab_size,
d_model=D_MODEL,
nhead=NUM_HEADS,
num_layers=NUM_LAYERS,
dropout=DROPOUT,
max_len=MAX_SEQ_LENGTH
).to(device)
# 4. Training Setup and Execution
dataset = TextDataset(training_data, current_tokenizer, MAX_SEQ_LENGTH)
data_loader = DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=True)
optimizer = Adam(current_model.parameters(), lr=LEARNING_RATE)
criterion = nn.CrossEntropyLoss(ignore_index=current_tokenizer.word_to_idx["<PAD>"])
print(f"\n[TRAINING] Starting {epochs} epochs with {len(dataset)} examples...")
train_model(current_model, data_loader, optimizer, criterion, device, epochs)
# 5. Persistence Update (Saves data if it was a retraining session)
if use_live_data:
# 5a. Update the base list to include the new data
initial_training_texts = training_data
# 5b. Save the combined data permanently to the CSV
save_data_to_csv(DATA_FILE, initial_training_texts)
# 5c. Clear the temporary queue
live_data_updates = []
print("[SYSTEM] Retraining complete. New knowledge acquired and **permanently saved**.")
def interactive_mode():
global live_data_updates, last_generated_text, last_user_prompt
# Check if the file exists before initial training
file_existed_before_run = os.path.exists(DATA_FILE)
global device
device = torch.device("cuda")
if not torch.cuda.is_available():
print("CRITICAL FAILURE: RTX NOT DETECTED! ABORTING EXPERIMENT!")
exit() # This kills the program rather than falling back to the i5!
# Run the initial, long training session
print(f"[SYSTEM] Using device: {device}")
initialize_or_retrain(initial_train=True, use_live_data=False, epochs=NUM_EPOCHS)
# IMPORTANT: If the file did not exist before this run (meaning default data was used),
# we force a save right now to write the 27 default sentences to the CSV file immediately.
if not file_existed_before_run:
print("\n[SYSTEM] CSV file was empty/missing. Forcing initial save of default knowledge...")
save_data_to_csv(DATA_FILE, initial_training_texts)
print("[SYSTEM] Default 27 sentences are now permanently written to training_data.csv.")
print("\n" + "=" * 60)
print("🤖 INTERACTIVE CHAT & LEARNING MODE 🤖")
print("1. Type a phrase to generate text (max 10 words).")
print("2. Use '!add [sentence]' to queue new training data.")
print("3. Use '!accept' to add the model's last **full** sentence to the training queue.")
print(f"4. Use '!retrain' to re-train the model on new data (runs for {INTERACTIVE_EPOCHS} epochs) **and save it**.")
print(f"5. Use '!refine' to re-train on existing data (runs for {INTERACTIVE_EPOCHS} epochs) **without saving.**")
print("6. Use '!penalty <value>' to regenerate with a different repetition penalty (higher = less repetition).")
print("7. Type 'quit' or 'exit' to stop.")
print("=" * 60)
while True:
try:
user_input = input("You: ")
if user_input.lower() in ['quit', 'exit']:
break
if user_input.lower().startswith('!add '):
sentence = user_input[5:].strip()
if sentence:
live_data_updates.append(sentence)
print(f"[SYSTEM] Added sentence to update queue: '{sentence}'")
print(f"[SYSTEM] Current update queue size: {len(live_data_updates)}. Type '!retrain' to apply and save changes.")
last_generated_text = None # Clear accepted text
last_user_prompt = None
continue
# --- !ACCEPT COMMAND ---
if user_input.lower().strip() == '!accept':
if last_generated_text and last_user_prompt:
# CRITICAL: Reconstruct the full sentence by joining prompt and output
full_sentence_parts = [last_user_prompt.strip(), last_generated_text.strip()]
sentence_to_add = " ".join(full_sentence_parts)
# Basic cleaning: ensure there aren't double spaces
sentence_to_add = " ".join(sentence_to_add.split())
if sentence_to_add and len(sentence_to_add.split()) > 4:
live_data_updates.append(sentence_to_add)
print(f"[SYSTEM] ACCEPTED: The full sentence '{sentence_to_add}' added to update queue.")
print(f"[SYSTEM] Current update queue size: {len(live_data_updates)}. Type '!retrain' to apply and save changes.")
last_generated_text = None # Clear after acceptance
last_user_prompt = None
else:
print("[SYSTEM] Cannot accept: The reconstructed sentence was too short or incomplete. Please use '!add [full sentence]' instead.")
else:
print("[SYSTEM] No text generated or prompt found. Generate text first.")
continue
# --- END !ACCEPT COMMAND ---
if user_input.lower() == '!retrain':
if not live_data_updates:
print("[SYSTEM] No new data to train on. Use '!add [sentence]' first.")
continue
print(f"\n[SYSTEM] RETRAINING MODEL ON NEW DATA ({INTERACTIVE_EPOCHS} EPOCHS)...")
initialize_or_retrain(initial_train=False, use_live_data=True, epochs=INTERACTIVE_EPOCHS)
last_generated_text = None # Clear the accepted text cache
last_user_prompt = None
continue
if user_input.lower() == '!refine':
print(f"\n[SYSTEM] REFINING MODEL ON EXISTING DATA ({INTERACTIVE_EPOCHS} EPOCHS)...")
initialize_or_retrain(initial_train=False, use_live_data=False, epochs=INTERACTIVE_EPOCHS)
print("[SYSTEM] Refinement complete. Knowledge deepened on existing data.")
continue
# --- !PENALTY COMMAND ---
if user_input.lower().startswith('!penalty '):
try:
penalty_value = float(user_input[9:].strip())
if penalty_value <= 0:
print("[SYSTEM] Penalty must be positive. Current penalty:", REPETITION_PENALTY)
continue
if not last_user_prompt:
print("[SYSTEM] No previous prompt to regenerate. Type a prompt first.")
continue
# Regenerate with new penalty
print(f"[SYSTEM] Regenerating with penalty={penalty_value}...")
generated_text = generate_text(
current_model,
current_tokenizer,
last_user_prompt,
MAX_SEQ_LENGTH,
device,
TOP_K,
penalty_value,
TEMPERATURE
)
print(f"Model: {generated_text}")
last_generated_text = generated_text
print("\n[HINT] If this full sentence is perfect, type '!accept' to add it to the training queue.")
except ValueError:
print(f"[SYSTEM] Invalid penalty value. Usage: !penalty <number>. Current penalty: {REPETITION_PENALTY}")
continue
# --- END !PENALTY COMMAND ---
if user_input.strip() and not user_input.lower().startswith(('!',)):
# Text generation logic
prompt = user_input.strip()
if len(prompt.split()) > MAX_SEQ_LENGTH - 1:
print(f"[SYSTEM] Prompt too long. Max {MAX_SEQ_LENGTH - 1} words supported.")
last_generated_text = None
last_user_prompt = None
continue
# 1. Store the prompt BEFORE generation
last_user_prompt = prompt
generated_text = generate_text(
current_model,
current_tokenizer,
prompt,
MAX_SEQ_LENGTH,
device,
TOP_K,
REPETITION_PENALTY,
TEMPERATURE
)
print(f"Model: {generated_text}")
# 2. Store the continuation AFTER generation
last_generated_text = generated_text
print("\n[HINT] If this full sentence is perfect, type '!accept' to add it to the training queue.")
except KeyboardInterrupt:
print("\nExiting interactive mode.")
break
except Exception as e:
print(f"An error occurred: {e}")
break
if __name__ == "__main__":
interactive_mode()