Download Python/ai thingy.py from Aobangaming/Aoban1.1A-Refined: direct link, hf CLI and curl.
- Browser
- Download file 28.2 kB
-
https://huggingface.co/Aobangaming/Aoban1.1A-Refined/resolve/b73f1278a1d944e12789b7ce45d5b361c788bd5f/Python/ai%20thingy.py
- Command line
-
hf download 'hf://Aobangaming/Aoban1.1A-Refined@b73f1278a1d944e12789b7ce45d5b361c788bd5f/Python/ai thingy.py'
-
curl -L -o 'ai thingy.py' https://huggingface.co/Aobangaming/Aoban1.1A-Refined/resolve/b73f1278a1d944e12789b7ce45d5b361c788bd5f/Python/ai%20thingy.py
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>"]]) | |
| 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() | |