bmdavis commited on
Commit
4888baa
·
verified ·
1 Parent(s): 3537e86

Create train_sentiment_model.py

Browse files
Files changed (1) hide show
  1. train_sentiment_model.py +59 -0
train_sentiment_model.py ADDED
@@ -0,0 +1,59 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from datasets import load_dataset
2
+ from transformers import (
3
+ AutoTokenizer,
4
+ AutoModelForSequenceClassification,
5
+ Trainer,
6
+ TrainingArguments,
7
+ )
8
+ import torch
9
+
10
+ # STEP 1: Load IMDb Dataset
11
+ dataset = load_dataset("imdb")
12
+
13
+ # STEP 2: Tokenize the Data
14
+ checkpoint = "distilbert-base-uncased"
15
+ tokenizer = AutoTokenizer.from_pretrained(checkpoint)
16
+
17
+ def preprocess(example):
18
+ return tokenizer(example["text"], truncation=True, padding="max_length", max_length=256)
19
+
20
+ tokenized = dataset.map(preprocess, batched=True)
21
+ tokenized = tokenized.remove_columns(["text"])
22
+ tokenized = tokenized.rename_column("label", "labels")
23
+ tokenized.set_format("torch")
24
+
25
+ # Use a smaller subset for quick training
26
+ train_dataset = tokenized["train"].shuffle(seed=42).select(range(2000))
27
+ val_dataset = tokenized["test"].shuffle(seed=42).select(range(500))
28
+
29
+ # STEP 3: Load Model
30
+ model = AutoModelForSequenceClassification.from_pretrained(checkpoint, num_labels=2)
31
+
32
+ # STEP 4: Define Training Arguments
33
+ training_args = TrainingArguments(
34
+ output_dir="./results",
35
+ evaluation_strategy="epoch",
36
+ save_strategy="epoch",
37
+ num_train_epochs=3,
38
+ per_device_train_batch_size=8,
39
+ per_device_eval_batch_size=8,
40
+ logging_dir="./logs",
41
+ logging_steps=50,
42
+ report_to="none"
43
+ )
44
+
45
+ # STEP 5: Train
46
+ trainer = Trainer(
47
+ model=model,
48
+ args=training_args,
49
+ train_dataset=train_dataset,
50
+ eval_dataset=val_dataset,
51
+ tokenizer=tokenizer,
52
+ )
53
+
54
+ trainer.train()
55
+
56
+ # STEP 6: Save Locally to Repo Folder
57
+ model.save_pretrained("./")
58
+ tokenizer.save_pretrained("./")
59
+ print("✅ Model and tokenizer saved locally!")