myAI / app.py
HedronCreeper's picture
Update app.py
92923aa verified
Raw
History Blame
2.44 kB
# app.py
import gradio as gr
from transformers import AutoTokenizer, AutoModelForCausalLM, Trainer, TrainingArguments
from datasets import load_dataset
import torch
import os
# -----------------------------
# 1️⃣ Model setup
# -----------------------------
MODEL_DIR = "model"
MODEL_NAME = "sshleifer/tiny-gpt2" # tiny GPT-2, CPU-friendly
# Load tokenizer & model
tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
model = AutoModelForCausalLM.from_pretrained(MODEL_NAME)
# Fix padding issue
tokenizer.pad_token = tokenizer.eos_token
# -----------------------------
# 2️⃣ Dataset setup
# -----------------------------
# Make sure you have 'data.txt' in the same folder as app.py
dataset = load_dataset("text", data_files="data.txt")
def tokenize(example):
return tokenizer(
example["text"],
truncation=True,
padding="max_length",
max_length=64 # small for CPU
)
tokenized_dataset = dataset.map(tokenize, batched=True)
# -----------------------------
# 3️⃣ Training setup
# -----------------------------
training_args = TrainingArguments(
output_dir=MODEL_DIR,
overwrite_output_dir=True,
per_device_train_batch_size=1, # CPU-friendly
num_train_epochs=1, # short test run
logging_steps=5,
save_steps=20,
save_total_limit=1
)
trainer = Trainer(
model=model,
args=training_args,
train_dataset=tokenized_dataset["train"]
)
# -----------------------------
# 4️⃣ Gradio interface
# -----------------------------
def train_model():
trainer.train()
return "✅ Training complete! Model saved to /model"
def generate_text(prompt):
inputs = tokenizer(prompt, return_tensors="pt", padding=True)
output = model.generate(**inputs, max_length=64, pad_token_id=tokenizer.eos_token_id)
return tokenizer.decode(output[0], skip_special_tokens=True)
with gr.Blocks() as demo:
gr.Markdown("# Tiny AI Training Demo")
with gr.Tab("Train Model"):
train_button = gr.Button("Train")
train_output = gr.Textbox(label="Logs")
train_button.click(train_model, outputs=train_output)
with gr.Tab("Generate Text"):
prompt_input = gr.Textbox(label="Prompt")
generate_button = gr.Button("Generate")
generate_output = gr.Textbox(label="Output")
generate_button.click(generate_text, inputs=prompt_input, outputs=generate_output)
demo.launch(share=True)