simonko912's picture
download
raw
11.3 kB
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.