BotX2.O / app.py
BotXT's picture
Update app.py
e838ec6 verified
Raw
History Blame Contribute Delete
3.35 kB
# app.py
import os
from datasets import load_dataset
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
Trainer,
TrainingArguments,
DataCollatorForLanguageModeling
)
import gradio as gr
import torch
# -------------------------
# 1. Load Dataset
# -------------------------
print("Loading dataset...")
dataset = load_dataset("json", data_files={
"train": [
"harmless-base/train.jsonl.gz",
"helpful-base/train.jsonl.gz",
"helpful-online/train.jsonl.gz",
"helpful-rejection-sampled/train.jsonl.gz"
],
"test": [
"harmless-base/test.jsonl.gz",
"helpful-base/test.jsonl.gz",
"helpful-online/test.jsonl.gz",
"helpful-rejection-sampled/test.jsonl.gz"
]
})
# -------------------------
# 2. Load Model & Tokenizer
# -------------------------
model_name = "distilgpt2"
print(f"Loading model: {model_name}")
tokenizer = AutoTokenizer.from_pretrained(model_name)
model = AutoModelForCausalLM.from_pretrained(model_name)
# Ensure tokenizer pads if needed
tokenizer.pad_token = tokenizer.eos_token
# -------------------------
# 3. Tokenize Dataset
# -------------------------
def tokenize_function(example):
# Combine 'input' and 'output' if they exist
text = example.get("input", "") + "\n" + example.get("output", "")
return tokenizer(text, truncation=True, padding="max_length", max_length=128)
print("Tokenizing dataset...")
tokenized_datasets = dataset.map(tokenize_function, batched=True)
# -------------------------
# 4. Data Collator
# -------------------------
data_collator = DataCollatorForLanguageModeling(
tokenizer=tokenizer, mlm=False
)
# -------------------------
# 5. Training Arguments
# -------------------------
training_args = TrainingArguments(
output_dir="./trained_model",
overwrite_output_dir=True,
num_train_epochs=1,
per_device_train_batch_size=8,
per_device_eval_batch_size=8,
save_steps=500,
save_total_limit=2,
logging_steps=50,
evaluation_strategy="steps",
eval_steps=200,
learning_rate=5e-5,
fp16=torch.cuda.is_available(),
push_to_hub=False
)
# -------------------------
# 6. Trainer
# -------------------------
trainer = Trainer(
model=model,
args=training_args,
train_dataset=tokenized_datasets["train"],
eval_dataset=tokenized_datasets["test"],
tokenizer=tokenizer,
data_collator=data_collator
)
# -------------------------
# 7. Train the Model
# -------------------------
print("Starting training...")
trainer.train()
print("Training complete! Saving model...")
trainer.save_model("./trained_model")
tokenizer.save_pretrained("./trained_model")
# -------------------------
# 8. Launch Gradio Interface
# -------------------------
print("Launching Gradio interface...")
def generate_response(prompt):
inputs = tokenizer(prompt, return_tensors="pt")
outputs = model.generate(
**inputs,
max_length=150,
do_sample=True,
top_p=0.95,
top_k=50
)
return tokenizer.decode(outputs[0], skip_special_tokens=True)
gr.Interface(
fn=generate_response,
inputs="text",
outputs="text",
title="Fine-tuned GPT-2 Chat",
description="Ask the fine-tuned GPT-2 model anything!"
).launch(server_name="0.0.0.0", server_port=int(os.environ.get("PORT", 7860)))