Skip to content

Commit 08680a8

Browse files
authored
Raise a clear error when ValidSplit gets an IterableDataset (#1151)
1 parent 2b54e3a commit 08680a8

3 files changed

Lines changed: 98 additions & 1 deletion

File tree

CHANGES.md

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,8 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
1616

1717
### Fixed
1818

19+
- `ValidSplit` now raises a clear error when it receives an `IterableDataset`, instead of an opaque `TypeError` about a missing length (#594)
20+
1921
## [1.4.0]
2022

2123
### Added

skorch/dataset.py

Lines changed: 10 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -312,7 +312,16 @@ def __call__(self, dataset, y=None, groups=None):
312312
raise bad_y_error
313313

314314
# pylint: disable=invalid-name
315-
len_dataset = get_len(dataset)
315+
try:
316+
len_dataset = get_len(dataset)
317+
except TypeError as exc:
318+
if not isinstance(dataset, torch.utils.data.IterableDataset):
319+
raise
320+
raise ValueError(
321+
"Cannot perform a CV split on an IterableDataset because it has "
322+
"no length. Set train_split=None to disable the internal "
323+
"validation split, or pass a train_split that supports "
324+
"IterableDataset.") from exc
316325
if y is not None:
317326
len_y = get_len(y)
318327
if len_dataset != len_y:

skorch/tests/test_net.py

Lines changed: 86 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2215,6 +2215,92 @@ def test_fit_with_dataset_stratified_without_explicit_y_raises(
22152215
msg = "Stratified CV requires explicitly passing a suitable y."
22162216
assert exc.value.args[0] == msg
22172217

2218+
@pytest.fixture
2219+
def iterable_dataset_cls(self):
2220+
class MyIterableDataset(torch.utils.data.IterableDataset):
2221+
def __init__(self, X, y):
2222+
super().__init__()
2223+
self.X = X
2224+
self.y = y
2225+
2226+
def __iter__(self):
2227+
return iter(zip(self.X, self.y))
2228+
2229+
return MyIterableDataset
2230+
2231+
def test_fit_with_iterable_dataset_and_train_split_raises(
2232+
self, net_cls, module_cls, iterable_dataset_cls, data):
2233+
from skorch.dataset import ValidSplit
2234+
2235+
net = net_cls(
2236+
module_cls,
2237+
max_epochs=1,
2238+
train_split=ValidSplit(stratified=False),
2239+
)
2240+
ds = iterable_dataset_cls(*data)
2241+
with pytest.raises(ValueError) as exc:
2242+
net.fit(ds, None)
2243+
2244+
msg = ("Cannot perform a CV split on an IterableDataset because it has "
2245+
"no length. Set train_split=None to disable the internal "
2246+
"validation split, or pass a train_split that supports "
2247+
"IterableDataset.")
2248+
assert exc.value.args[0] == msg
2249+
assert isinstance(exc.value.__cause__, TypeError)
2250+
2251+
def test_fit_with_iterable_dataset_no_train_split(
2252+
self, net_cls, module_cls, iterable_dataset_cls, data):
2253+
net = net_cls(module_cls, max_epochs=1, train_split=None)
2254+
ds = iterable_dataset_cls(*data)
2255+
net.fit(ds, None) # does not raise
2256+
2257+
assert 'train_loss' in net.history[-1]
2258+
2259+
@pytest.fixture
2260+
def sized_iterable_dataset_cls(self):
2261+
class MySizedIterableDataset(torch.utils.data.IterableDataset):
2262+
def __init__(self, X, y):
2263+
super().__init__()
2264+
self.X = X
2265+
self.y = y
2266+
2267+
def __iter__(self):
2268+
return iter(zip(self.X, self.y))
2269+
2270+
def __len__(self):
2271+
return len(self.y)
2272+
2273+
def __getitem__(self, i):
2274+
return self.X[i], self.y[i]
2275+
2276+
return MySizedIterableDataset
2277+
2278+
def test_fit_with_sized_iterable_dataset_and_train_split(
2279+
self, net_cls, module_cls, sized_iterable_dataset_cls, data):
2280+
from skorch.dataset import ValidSplit
2281+
2282+
net = net_cls(
2283+
module_cls,
2284+
max_epochs=1,
2285+
train_split=ValidSplit(stratified=False),
2286+
)
2287+
ds = sized_iterable_dataset_cls(*data)
2288+
net.fit(ds, None) # does not raise
2289+
2290+
assert 'valid_loss' in net.history[-1]
2291+
2292+
def test_fit_with_iterable_dataset_and_custom_train_split(
2293+
self, net_cls, module_cls, iterable_dataset_cls, data):
2294+
# a train_split that never indexes the dataset must keep working
2295+
def train_split(dataset, **kwargs):
2296+
return dataset, dataset
2297+
2298+
net = net_cls(module_cls, max_epochs=1, train_split=train_split)
2299+
ds = iterable_dataset_cls(*data)
2300+
net.fit(ds, None) # does not raise
2301+
2302+
assert 'valid_loss' in net.history[-1]
2303+
22182304
@pytest.fixture
22192305
def dataset_1_item(self):
22202306
class Dataset(torch.utils.data.Dataset):

0 commit comments

Comments
 (0)