File size: 3,353 Bytes
e838ec6
 
 
b2b1f43
e838ec6
 
 
 
 
 
 
 
c14ddf2
b2b1f43
e838ec6
 
 
b2b1f43
e838ec6
 
 
 
 
 
 
 
 
 
 
 
 
 
b2b1f43
e838ec6
 
 
b2b1f43
e838ec6
b2b1f43
c14ddf2
b2b1f43
e838ec6
 
b2b1f43
e838ec6
 
 
b2b1f43
e838ec6
 
 
b2b1f43
 
e838ec6
b2b1f43
e838ec6
 
 
 
 
 
 
 
 
 
b2b1f43
e838ec6
 
c14ddf2
e838ec6
 
 
 
 
 
 
c14ddf2
 
e838ec6
b2b1f43
 
e838ec6
 
 
b2b1f43
 
 
e838ec6
 
 
 
b2b1f43
 
e838ec6
 
 
c14ddf2
b2b1f43
e838ec6
 
c14ddf2
b2b1f43
e838ec6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4935bcb
e838ec6
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
# 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)))