Skip to content

Commit 9203e6a

Browse files
author
Federico Errica
committed
removed sampler since it is not used anymore
1 parent e6403db commit 9203e6a

5 files changed

Lines changed: 5 additions & 69 deletions

File tree

docs/source/mlwiz.data.rst

Lines changed: 0 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -19,15 +19,6 @@ data.provider
1919
:undoc-members:
2020
:show-inheritance:
2121

22-
data.sampler
23-
-------------------------
24-
25-
.. automodule:: mlwiz.data.sampler
26-
:members:
27-
:private-members:
28-
:undoc-members:
29-
:show-inheritance:
30-
3122
data.splitter
3223
--------------------------
3324

mlwiz/data/provider.py

Lines changed: 2 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,6 @@
1717

1818
import mlwiz.data.dataset
1919
from mlwiz.data.dataset import DatasetInterface
20-
from mlwiz.data.sampler import RandomSampler
2120
from mlwiz.data.splitter import Splitter, SingleGraphSplitter
2221
from mlwiz.data.util import load_dataset, single_graph_collate
2322

@@ -133,10 +132,6 @@ class DataProvider:
133132
special, but here is where the i-th element of a dataset could be
134133
pre-processed before constructing the mini-batches.
135134
136-
IMPORTANT: if the dataset is to be shuffled, you MUST use a
137-
:class:`mlwiz.data.sampler.RandomSampler` object to determine the
138-
permutation.
139-
140135
Args:
141136
storage_folder (str): the path of the root folder in which data is stored
142137
splits_filepath (str): the filepath of the splits. with additional
@@ -369,9 +364,8 @@ def _get_loader(
369364
dataset, sampler=sampler, batch_size=batch_size, **kwargs
370365
)
371366
elif shuffle is True:
372-
sampler = RandomSampler(dataset)
373367
dataloader = self.data_loader_class(
374-
dataset, sampler=sampler, batch_size=batch_size, **kwargs
368+
dataset, shuffle=True, batch_size=batch_size, **kwargs
375369
)
376370
else:
377371
dataloader = self.data_loader_class(
@@ -725,10 +719,9 @@ def _get_loader(
725719
kwargs.update(self.data_loader_args)
726720

727721
if shuffle is True:
728-
sampler = RandomSampler(dataset)
729722
dataloader = self.data_loader_class(
730723
dataset,
731-
sampler=sampler,
724+
shuffle=True,
732725
batch_size=batch_size,
733726
collate_fn=single_graph_collate,
734727
**kwargs,

mlwiz/data/sampler.py

Lines changed: 0 additions & 47 deletions
This file was deleted.

tests/data/test_data_provider_additional.py

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@
1111

1212
import pytest
1313
import torch
14-
from torch.utils.data import DataLoader, SequentialSampler
14+
from torch.utils.data import DataLoader, RandomSampler, SequentialSampler
1515

1616
import mlwiz.data.provider as provider_mod
1717
from mlwiz.data.provider import (
@@ -20,7 +20,6 @@
2020
SubsetTrainEval,
2121
_iterable_worker_init_fn,
2222
)
23-
from mlwiz.data.sampler import RandomSampler
2423
from mlwiz.data.splitter import Splitter
2524

2625

@@ -191,7 +190,7 @@ def test_data_provider_requires_seed_and_fold_ids():
191190

192191

193192
def test_data_provider_get_loader_uses_random_sampler_when_shuffling():
194-
"""Ensure shuffle=True uses MLWiz ``RandomSampler`` (with stored permutation)."""
193+
"""Ensure shuffle=True uses a random sampler."""
195194
provider = DataProvider(
196195
storage_folder="unused",
197196
splits_filepath="unused",

tests/training/test_engine_helpers.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -109,7 +109,7 @@ def _raise_termination():
109109
engine._check_termination()
110110

111111

112-
def test_infer_sequential_sampler_requires_no_permutation_and_sets_main_keys(
112+
def test_infer_sets_main_keys(
113113
tmp_path,
114114
):
115115
"""infer() should set MAIN_* keys for downstream consumers."""

0 commit comments

Comments
 (0)