File size: 2,025 Bytes
166edf7 0c2ae95 166edf7 0c2ae95 166edf7 0c2ae95 166edf7 0c2ae95 166edf7 0c2ae95 166edf7 2c9e8bc 166edf7 | 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 | import os
from torch.utils.data import DataLoader
from lightning import LightningDataModule
from .mixed_dataset import MixedDataset
class MixedDataModule(LightningDataModule):
def __init__(
self, bert_model, dataset_path, tool_capacity, batch_size, num_workers, seed
):
super().__init__()
self.bert_model = bert_model
self.dataset_path = dataset_path
self.tool_capacity = tool_capacity
self.batch_size = batch_size
self.num_workers = num_workers
self.seed = seed
def setup(self, stage=None):
if stage == "fit":
self.train_dataset = MixedDataset(
self.bert_model,
"train",
os.path.join(self.dataset_path, "train.json"),
self.tool_capacity,
seed=self.seed,
)
self.val_dataset = MixedDataset(
self.bert_model,
"test",
os.path.join(self.dataset_path, "test.json"),
self.tool_capacity,
seed=self.seed,
)
elif stage == "test":
self.test_dataset = MixedDataset(
self.bert_model,
"test",
os.path.join(self.dataset_path, "test.json"),
self.tool_capacity,
seed=self.seed,
)
def train_dataloader(self):
return DataLoader(
self.train_dataset,
batch_size=self.batch_size,
shuffle=True,
num_workers=self.num_workers,
drop_last=True,
)
def val_dataloader(self):
return DataLoader(
self.val_dataset,
batch_size=self.batch_size,
shuffle=False,
num_workers=self.num_workers,
)
def test_dataloader(self):
return DataLoader(
self.test_dataset,
batch_size=self.batch_size,
shuffle=False,
num_workers=self.num_workers,
)
|