Buckets:
| import json | |
| import time | |
| from datetime import datetime | |
| import torch | |
| import os | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| print(f"Using device: {device}") | |
| from datasets import Dataset | |
| from transformers import ( | |
| Trainer, | |
| TrainingArguments, | |
| DataCollatorForLanguageModeling, | |
| ) | |
| # ========================= | |
| # Architecture imports | |
| # ========================= | |
| from transformers import GPT2Config, GPT2LMHeadModel, GPT2TokenizerFast | |
| from transformers import LlamaConfig, LlamaForCausalLM, LlamaTokenizer | |
| # ========================= | |
| # Config | |
| # ========================= | |
| OUTPUT_DIR = input("Enter the path where you want to save the model: ").strip() | |
| epoch_num = int(input("Number of epoches: ").strip()) | |
| # ========================= | |
| # Training time window | |
| # (solar energy schedule) | |
| # ========================= | |
| TRAINING_WINDOW_ENABLED = True # Set to False to train at any time | |
| TRAINING_WINDOW_START = 5 # Hour (24h): training is allowed from this hour | |
| TRAINING_WINDOW_END = 18 # Hour (24h): training pauses at this hour | |
| SIZE_MODES = { | |
| "banana": dict(n_layer=2, n_head=2, n_embd=32, hidden_size=32, intermediate_size=128), | |
| "nano": dict(n_layer=2, n_head=2, n_embd=64, hidden_size=64, intermediate_size=256), | |
| "small": dict(n_layer=4, n_head=4, n_embd=128, hidden_size=128, intermediate_size=512), | |
| "medium": dict(n_layer=6, n_head=6, n_embd=384, hidden_size=384, intermediate_size=1536), | |
| "large": dict(n_layer=12, n_head=12, n_embd=768, hidden_size=768, intermediate_size=3072), | |
| "larger": dict(n_layer=24, n_head=16, n_embd=1024, hidden_size=1024, intermediate_size=4096), | |
| } | |
| MAX_LENGTH = 4096 | |
| # ========================= | |
| # Time window helpers | |
| # ========================= | |
| def is_within_training_window(): | |
| """Return True if the current hour is within the allowed training window.""" | |
| if not TRAINING_WINDOW_ENABLED: | |
| return True | |
| hour = datetime.now().hour | |
| return TRAINING_WINDOW_START <= hour < TRAINING_WINDOW_END | |
| def wait_for_training_window(): | |
| """Block, sleeping 60 s at a time, until the training window opens again.""" | |
| if is_within_training_window(): | |
| return | |
| now = datetime.now() | |
| print(f"\n⏸ Outside training window ({now.strftime('%H:%M')}). " | |
| f"Pausing until {TRAINING_WINDOW_START:02d}:00 ...") | |
| while not is_within_training_window(): | |
| time.sleep(60) | |
| print(f"▶ Training window open ({datetime.now().strftime('%H:%M')}). Resuming ...") | |
| # Trainer callback — checks the clock after every optimizer step | |
| from transformers import TrainerCallback | |
| class TimeWindowCallback(TrainerCallback): | |
| def on_step_end(self, args, state, control, **kwargs): | |
| if TRAINING_WINDOW_ENABLED and not is_within_training_window(): | |
| wait_for_training_window() | |
| # ========================= | |
| # Dataset builders | |
| # ========================= | |
| def build_dataset_from_txt(path): | |
| def gen(): | |
| buffer = [] | |
| with open(path, "r", encoding="utf-8") as f: | |
| for line in f: | |
| if line.strip() == "": | |
| if buffer: | |
| yield {"text": "\n".join(buffer)} | |
| buffer = [] | |
| else: | |
| buffer.append(line.rstrip("\n")) | |
| if buffer: | |
| yield {"text": "\n".join(buffer)} | |
| return Dataset.from_generator(gen) | |
| def build_dataset_from_chat_json(path): | |
| # full JSON array (old) | |
| with open(path, "r", encoding="utf-8") as f: | |
| data = json.load(f) | |
| samples = [] | |
| for convo in data: | |
| lines = [] | |
| for msg in convo["messages"]: | |
| role = msg["role"].capitalize() | |
| lines.append(f"{role}: {msg['content']}") # use 'content' for OASST JSONL | |
| samples.append("\n".join(lines)) | |
| return Dataset.from_dict({"text": samples}) | |
| def build_dataset_from_chat_jsonl(path): | |
| samples = [] | |
| with open(path, "r", encoding="utf-8") as f: | |
| for line in f: | |
| convo = json.loads(line) | |
| lines = [] | |
| for msg in convo["messages"]: | |
| role = msg["role"].capitalize() | |
| # include input if it exists | |
| if "input" in msg and msg["input"].strip(): | |
| lines.append(f"{role} [Input: {msg['input']}]: {msg['content']}") | |
| else: | |
| lines.append(f"{role}: {msg['content']}") | |
| samples.append("\n".join(lines)) | |
| return Dataset.from_dict({"text": samples}) | |
| # ========================= | |
| # Preprocess factory | |
| # ========================= | |
| def preprocess_factory(tokenizer): | |
| def preprocess(example): | |
| enc = tokenizer( | |
| example["text"], | |
| truncation=True, | |
| padding="max_length", | |
| max_length=MAX_LENGTH, | |
| ) | |
| enc["labels"] = enc["input_ids"].copy() | |
| return enc | |
| return preprocess | |
| # ========================= | |
| # Train from scratch | |
| # ========================= | |
| def train_scratch(): | |
| print("\nDataset type:") | |
| print("1) Plain .txt (recommended)") | |
| print("2) Chat JSON (full JSON array)") | |
| print("3) Chat JSONL (line-delimited, OpenAssistant-style)") | |
| choice = input("Choice: ").strip() | |
| if choice == "1": | |
| path = input("Path to .txt file: ").strip() | |
| dataset = build_dataset_from_txt(path) | |
| elif choice == "2": | |
| path = input("Path to chat JSON: ").strip() | |
| dataset = build_dataset_from_chat_json(path) | |
| elif choice == "3": | |
| path = input("Path to chat JSONL: ").strip() | |
| dataset = build_dataset_from_chat_jsonl(path) | |
| else: | |
| print("❌ Invalid choice") | |
| return | |
| # ========================= | |
| # Resume from checkpoint? | |
| # ======================== | |
| resume_from_checkpoint = None | |
| if os.path.exists(OUTPUT_DIR) and os.listdir(OUTPUT_DIR): | |
| print("\n📂 Existing checkpoint(s) found") | |
| resume_choice = input("Resume from checkpoint? (y/n): ").strip().lower() | |
| if resume_choice == "y": | |
| resume_from_checkpoint = input( | |
| "Enter checkpoint path (e.g. output_dir/checkpoint-2000): " | |
| ).strip() | |
| print("✅ Resuming from:", resume_from_checkpoint) | |
| # ========================= | |
| # Architecture choice | |
| # ========================= | |
| print("\nArchitecture:") | |
| print("1) GPT-2") | |
| print("2) LLaMA (tiny from scratch)") | |
| arch_choice = input("Select architecture: ").strip() | |
| print("\nModel size:") | |
| for k in SIZE_MODES: | |
| print("-", k) | |
| size = input("Size: ").strip() | |
| if size not in SIZE_MODES: | |
| print("❌ Invalid size") | |
| return | |
| # ========================= | |
| # GPT-2 | |
| # ========================= | |
| if arch_choice == "1": | |
| print(f"\n🧠 Training GPT-2 from scratch ({size})") | |
| tokenizer = GPT2TokenizerFast.from_pretrained("gpt2") | |
| tokenizer.pad_token = tokenizer.eos_token | |
| torch.backends.cuda.enable_flash_sdp(True) | |
| config = GPT2Config( | |
| vocab_size=tokenizer.vocab_size, | |
| n_positions=MAX_LENGTH, | |
| n_ctx=MAX_LENGTH, | |
| **SIZE_MODES[size] | |
| ) | |
| model = GPT2LMHeadModel(config).to(device) | |
| # ========================= | |
| # LLaMA (tiny) | |
| # ========================= | |
| elif arch_choice == "2": | |
| print(f"\n🧠 Training Tiny LLaMA from scratch ({size})") | |
| tokenizer = LlamaTokenizer.from_pretrained("huggyllama/llama-7b") | |
| tokenizer.pad_token = tokenizer.eos_token | |
| config = LlamaConfig( | |
| vocab_size=tokenizer.vocab_size, | |
| max_position_embeddings=MAX_LENGTH, | |
| num_attention_heads=SIZE_MODES[size]["n_head"], | |
| num_hidden_layers=SIZE_MODES[size]["n_layer"], | |
| hidden_size=SIZE_MODES[size]["hidden_size"], | |
| intermediate_size=SIZE_MODES[size]["intermediate_size"], | |
| attn_implementation="flash_attention_2", | |
| ) | |
| model = LlamaForCausalLM(config).to(device) | |
| else: | |
| print("❌ Invalid architecture") | |
| return | |
| # ========================= | |
| # Preprocess dataset | |
| # ========================= | |
| dataset = dataset.map( | |
| preprocess_factory(tokenizer), | |
| batched=True, | |
| num_proc=48, # ← explicitly set this high | |
| remove_columns=["text"], | |
| ) | |
| collator = DataCollatorForLanguageModeling( | |
| tokenizer=tokenizer, | |
| mlm=False | |
| ) | |
| args = TrainingArguments( | |
| output_dir=OUTPUT_DIR, | |
| num_train_epochs=epoch_num, | |
| per_device_train_batch_size=6, # or 6 | |
| gradient_accumulation_steps=4, # 4*3=12 effective batch | |
| gradient_checkpointing=True, | |
| fp16=True, | |
| save_steps=2000, | |
| save_total_limit=5, | |
| learning_rate=3e-4, | |
| logging_steps=50, | |
| save_strategy="steps", | |
| report_to="none", | |
| resume_from_checkpoint=resume_from_checkpoint, | |
| ) | |
| # Block here if training is started outside the allowed window | |
| wait_for_training_window() | |
| trainer = Trainer( | |
| model=model, | |
| args=args, | |
| train_dataset=dataset, | |
| data_collator=collator, | |
| callbacks=[TimeWindowCallback()], | |
| ) | |
| trainer.train(resume_from_checkpoint=resume_from_checkpoint) | |
| model.save_pretrained(OUTPUT_DIR) | |
| tokenizer.save_pretrained(OUTPUT_DIR) | |
| print("\n✅ Training complete") | |
| # ========================= | |
| # Run (inference) | |
| # ========================= | |
| def run(): | |
| # Try GPT-2 first, fallback to LLaMA | |
| try: | |
| tokenizer = GPT2TokenizerFast.from_pretrained(OUTPUT_DIR) | |
| tokenizer.pad_token = tokenizer.eos_token | |
| model = GPT2LMHeadModel.from_pretrained(OUTPUT_DIR) | |
| arch = "GPT-2" | |
| except: | |
| tokenizer = LlamaTokenizer.from_pretrained(OUTPUT_DIR) | |
| tokenizer.pad_token = tokenizer.eos_token | |
| model = LlamaForCausalLM.from_pretrained(OUTPUT_DIR) | |
| arch = "LLaMA" | |
| model = model.to(device) | |
| model.eval() | |
| print(f"\nRunning {arch} model. Type 'exit' to quit.") | |
| print("User: hello\nAssistant:") | |
| while True: | |
| user_prompt = input("\n> ") | |
| if user_prompt.lower() == "exit": | |
| break | |
| # Optional: if you want the assistant to "see" a fixed input | |
| assistant_input = input("Assistant input (optional, leave empty if none): ").strip() | |
| if assistant_input: | |
| prompt = f"User: {user_prompt}\nAssistant [Input: {assistant_input}]:" | |
| else: | |
| prompt = f"User: {user_prompt}\nAssistant:" | |
| inputs = tokenizer(prompt, return_tensors="pt").to(device) | |
| with torch.no_grad(): | |
| out = model.generate( | |
| **inputs, | |
| max_new_tokens=80, | |
| do_sample=True, | |
| temperature=0.8, | |
| top_p=0.95, | |
| ) | |
| print(tokenizer.decode(out[0], skip_special_tokens=True)) | |
| # ========================= | |
| # Menu | |
| # ========================= | |
| def main(): | |
| print(""" | |
| ============================= | |
| TRAIN GPT-2 | |
| OR TINY LLaMA | |
| ============================= | |
| 1) Train | |
| 2) Run | |
| 3) Exit | |
| """) | |
| c = input("Select: ").strip() | |
| if c == "1": | |
| train_scratch() | |
| elif c == "2": | |
| run() | |
| else: | |
| print("👋 Bye") | |
| if __name__ == "__main__": | |
| main() | |
Xet Storage Details
- Size:
- 11.3 kB
- Xet hash:
- 7d7706a523b4cb26f46a4803e879be0617ca4bb648c54d5a1b62e7e556891ff9
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.