@@ -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