-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdata_module.py
More file actions
executable file
·100 lines (92 loc) · 3.47 KB
/
Copy pathdata_module.py
File metadata and controls
executable file
·100 lines (92 loc) · 3.47 KB
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
from typing import Optional
import pandas as pd
import torch
from torch.utils.data import Dataset, DataLoader
from pytorch_lightning import LightningDataModule
from transformers import AutoTokenizer
class TextDataset(Dataset):
def __init__(self, texts, labels, tokenizer, max_length=128):
self.texts = texts
self.labels = labels
self.tokenizer = tokenizer
self.max_length = max_length
def __len__(self):
return len(self.texts)
def __getitem__(self, idx):
text = str(self.texts[idx])
label = self.labels[idx]
encoding = self.tokenizer.encode_plus(
text,
add_special_tokens=True,
max_length=self.max_length,
return_token_type_ids=False,
padding='max_length',
truncation=True,
return_attention_mask=True,
return_tensors='pt',
)
return {
'input_ids': encoding['input_ids'].squeeze(0),
'attention_mask': encoding['attention_mask'].squeeze(0),
'labels': torch.tensor(label, dtype=torch.long)
}
class TextClassificationDataModule(LightningDataModule):
def __init__(self, train_df=None, val_df=None, test_df=None,
tokenizer_name='answerdotai/ModernBERT-base', batch_size=16, max_length=128, num_workers=4):
super().__init__()
self.train_df = train_df
self.val_df = val_df
self.test_df = test_df
self.batch_size = batch_size
self.max_length = max_length
self.num_workers = num_workers
self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_name)
def setup(self, stage: Optional[str] = None):
if stage == 'fit' or stage is None:
if self.train_df is not None:
self.train_dataset = TextDataset(
texts=self.train_df['text'].values,
labels=self.train_df['label'].values,
tokenizer=self.tokenizer,
max_length=self.max_length
)
if self.val_df is not None:
self.val_dataset = TextDataset(
texts=self.val_df['text'].values,
labels=self.val_df['label'].values,
tokenizer=self.tokenizer,
max_length=self.max_length
)
if stage == 'test' or stage is None:
if self.test_df is not None:
self.test_dataset = TextDataset(
texts=self.test_df['text'].values,
labels=self.test_df['label'].values,
tokenizer=self.tokenizer,
max_length=self.max_length
)
def train_dataloader(self):
return DataLoader(
self.train_dataset,
batch_size=self.batch_size,
shuffle=True,
num_workers=self.num_workers,
pin_memory=True,
persistent_workers=self.num_workers > 0
)
def val_dataloader(self):
return DataLoader(
self.val_dataset,
batch_size=self.batch_size,
num_workers=self.num_workers,
pin_memory=True,
persistent_workers=self.num_workers > 0,
)
def test_dataloader(self):
return DataLoader(
self.test_dataset,
batch_size=self.batch_size,
num_workers=self.num_workers,
pin_memory=True,
persistent_workers=self.num_workers > 0,
)