| |
|
|
| import os |
| from datasets import load_dataset |
| from transformers import ( |
| AutoModelForCausalLM, |
| AutoTokenizer, |
| Trainer, |
| TrainingArguments, |
| DataCollatorForLanguageModeling |
| ) |
| import gradio as gr |
| import torch |
|
|
| |
| |
| |
| 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" |
| ] |
| }) |
|
|
| |
| |
| |
| model_name = "distilgpt2" |
| print(f"Loading model: {model_name}") |
| tokenizer = AutoTokenizer.from_pretrained(model_name) |
| model = AutoModelForCausalLM.from_pretrained(model_name) |
|
|
| |
| tokenizer.pad_token = tokenizer.eos_token |
|
|
| |
| |
| |
| def tokenize_function(example): |
| |
| 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) |
|
|
| |
| |
| |
| data_collator = DataCollatorForLanguageModeling( |
| tokenizer=tokenizer, mlm=False |
| ) |
|
|
| |
| |
| |
| 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 |
| ) |
|
|
| |
| |
| |
| trainer = Trainer( |
| model=model, |
| args=training_args, |
| train_dataset=tokenized_datasets["train"], |
| eval_dataset=tokenized_datasets["test"], |
| tokenizer=tokenizer, |
| data_collator=data_collator |
| ) |
|
|
| |
| |
| |
| print("Starting training...") |
| trainer.train() |
| print("Training complete! Saving model...") |
| trainer.save_model("./trained_model") |
| tokenizer.save_pretrained("./trained_model") |
|
|
| |
| |
| |
| 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))) |
|
|