Package mvpa :: Package tests :: Module test_splitter
[hide private]
[frames] | no frames]

Source Code for Module mvpa.tests.test_splitter

  1  # emacs: -*- mode: python; py-indent-offset: 4; indent-tabs-mode: nil -*- 
  2  # vi: set ft=python sts=4 ts=4 sw=4 et: 
  3  ### ### ### ### ### ### ### ### ### ### ### ### ### ### ### ### ### ### ### ## 
  4  # 
  5  #   See COPYING file distributed along with the PyMVPA package for the 
  6  #   copyright and license terms. 
  7  # 
  8  ### ### ### ### ### ### ### ### ### ### ### ### ### ### ### ### ### ### ### ## 
  9  """Unit tests for PyMVPA pattern handling""" 
 10   
 11  from mvpa.datasets.masked import MaskedDataset 
 12  from mvpa.datasets.splitters import NFoldSplitter, OddEvenSplitter, \ 
 13                                     NoneSplitter, HalfSplitter, \ 
 14                                     CustomSplitter, NGroupSplitter 
 15  import unittest 
 16  import numpy as N 
 17   
 18   
19 -class SplitterTests(unittest.TestCase):
20
21 - def setUp(self):
22 self.data = \ 23 MaskedDataset(samples=N.random.normal(size=(100,10)), 24 labels=[ i%4 for i in range(100) ], 25 chunks=[ i/10 for i in range(100)])
26 27
28 - def testSimplestCVPatGen(self):
29 # create the generator 30 nfs = NFoldSplitter(cvtype=1) 31 32 # now get the xval pattern sets One-Fold CV) 33 xvpat = [ (train, test) for (train,test) in nfs(self.data) ] 34 35 self.failUnless( len(xvpat) == 10 ) 36 37 for i,p in enumerate(xvpat): 38 self.failUnless( len(p) == 2 ) 39 self.failUnless( p[0].nsamples == 90 ) 40 self.failUnless( p[1].nsamples == 10 ) 41 self.failUnless( p[1].chunks[0] == i )
42 43
44 - def testOddEvenSplit(self):
45 oes = OddEvenSplitter() 46 47 splits = [ (train, test) for (train, test) in oes(self.data) ] 48 49 self.failUnless(len(splits) == 2) 50 51 for i,p in enumerate(splits): 52 self.failUnless( len(p) == 2 ) 53 self.failUnless( p[0].nsamples == 50 ) 54 self.failUnless( p[1].nsamples == 50 ) 55 56 self.failUnless((splits[0][1].uniquechunks == [1, 3, 5, 7, 9]).all()) 57 self.failUnless((splits[0][0].uniquechunks == [0, 2, 4, 6, 8]).all()) 58 self.failUnless((splits[1][0].uniquechunks == [1, 3, 5, 7, 9]).all()) 59 self.failUnless((splits[1][1].uniquechunks == [0, 2, 4, 6, 8]).all()) 60 61 # check if it works on pure odd and even chunk ids 62 moresplits = [ (train, test) for (train, test) in oes(splits[0][0])] 63 64 for split in moresplits: 65 self.failUnless(split[0] != None) 66 self.failUnless(split[1] != None)
67 68
69 - def testHalfSplit(self):
70 hs = HalfSplitter() 71 72 splits = [ (train, test) for (train, test) in hs(self.data) ] 73 74 self.failUnless(len(splits) == 2) 75 76 for i,p in enumerate(splits): 77 self.failUnless( len(p) == 2 ) 78 self.failUnless( p[0].nsamples == 50 ) 79 self.failUnless( p[1].nsamples == 50 ) 80 81 self.failUnless((splits[0][1].uniquechunks == [0, 1, 2, 3, 4]).all()) 82 self.failUnless((splits[0][0].uniquechunks == [5, 6, 7, 8, 9]).all()) 83 self.failUnless((splits[1][1].uniquechunks == [5, 6, 7, 8, 9]).all()) 84 self.failUnless((splits[1][0].uniquechunks == [0, 1, 2, 3, 4]).all()) 85 86 # check if it works on pure odd and even chunk ids 87 moresplits = [ (train, test) for (train, test) in hs(splits[0][0])] 88 89 for split in moresplits: 90 self.failUnless(split[0] != None) 91 self.failUnless(split[1] != None)
92
93 - def testNGroupSplit(self):
94 # Test 2 groups like HalfSplitter first 95 hs = NGroupSplitter(2) 96 splits = [ (train, test) for (train, test) in hs(self.data) ] 97 98 self.failUnless(len(splits) == 2) 99 100 for i,p in enumerate(splits): 101 self.failUnless( len(p) == 2 ) 102 self.failUnless( p[0].nsamples == 50 ) 103 self.failUnless( p[1].nsamples == 50 ) 104 105 self.failUnless((splits[0][1].uniquechunks == [0, 1, 2, 3, 4]).all()) 106 self.failUnless((splits[0][0].uniquechunks == [5, 6, 7, 8, 9]).all()) 107 self.failUnless((splits[1][1].uniquechunks == [5, 6, 7, 8, 9]).all()) 108 self.failUnless((splits[1][0].uniquechunks == [0, 1, 2, 3, 4]).all()) 109 110 # check if it works on pure odd and even chunk ids 111 moresplits = [ (train, test) for (train, test) in hs(splits[0][0])] 112 113 for split in moresplits: 114 self.failUnless(split[0] != None) 115 self.failUnless(split[1] != None) 116 117 # now test more groups 118 s5 = NGroupSplitter(5) 119 120 # get the splits 121 splits = [ (train, test) for (train, test) in s5(self.data) ] 122 123 # must have 10 splits 124 self.failUnless(len(splits) == 5) 125 126 # check split content 127 self.failUnless((splits[0][1].uniquechunks == [0, 1]).all()) 128 self.failUnless((splits[0][0].uniquechunks == [2, 3, 4, 5, 6, 7, 8, 9]).all()) 129 self.failUnless((splits[1][1].uniquechunks == [2, 3]).all()) 130 self.failUnless((splits[1][0].uniquechunks == [0, 1, 4, 5, 6, 7, 8, 9]).all()) 131 # ... 132 self.failUnless((splits[4][1].uniquechunks == [8, 9]).all()) 133 self.failUnless((splits[4][0].uniquechunks == [0, 1, 2, 3, 4, 5, 6, 7]).all()) 134 135 # Test for too many groups 136 def splitcall(spl, dat): 137 return [ (train, test) for (train, test) in spl(dat) ]
138 s20 = NGroupSplitter(20) 139 self.assertRaises(ValueError,splitcall,s20,self.data)
140
141 - def testCustomSplit(self):
142 #simulate half splitter 143 hs = CustomSplitter([(None,[0,1,2,3,4]),(None,[5,6,7,8,9])]) 144 splits = list(hs(self.data)) 145 self.failUnless(len(splits) == 2) 146 147 for i,p in enumerate(splits): 148 self.failUnless( len(p) == 2 ) 149 self.failUnless( p[0].nsamples == 50 ) 150 self.failUnless( p[1].nsamples == 50 ) 151 152 self.failUnless((splits[0][1].uniquechunks == [0, 1, 2, 3, 4]).all()) 153 self.failUnless((splits[0][0].uniquechunks == [5, 6, 7, 8, 9]).all()) 154 self.failUnless((splits[1][1].uniquechunks == [5, 6, 7, 8, 9]).all()) 155 self.failUnless((splits[1][0].uniquechunks == [0, 1, 2, 3, 4]).all()) 156 157 158 # check fully customized split with working and validation set specified 159 cs = CustomSplitter([([0,3,4],[5,9])]) 160 splits = list(cs(self.data)) 161 self.failUnless(len(splits) == 1) 162 163 for i,p in enumerate(splits): 164 self.failUnless( len(p) == 2 ) 165 self.failUnless( p[0].nsamples == 30 ) 166 self.failUnless( p[1].nsamples == 20 ) 167 168 self.failUnless((splits[0][1].uniquechunks == [5, 9]).all()) 169 self.failUnless((splits[0][0].uniquechunks == [0, 3, 4]).all()) 170 171 # full test with additional sampling and 3 datasets per split 172 cs = CustomSplitter([([0,3,4],[5,9],[2])], 173 nperlabel=[3,4,1], 174 nrunspersplit=3) 175 splits = list(cs(self.data)) 176 self.failUnless(len(splits) == 3) 177 178 for i,p in enumerate(splits): 179 self.failUnless( len(p) == 3 ) 180 self.failUnless( p[0].nsamples == 12 ) 181 self.failUnless( p[1].nsamples == 16 ) 182 self.failUnless( p[2].nsamples == 4 ) 183 184 # lets test selection of samples by ratio and combined with 185 # other ways 186 cs = CustomSplitter([([0,3,4],[5,9],[2])], 187 nperlabel=[[0.3, 0.6, 1.0, 0.5], 188 0.5, 189 'all'], 190 nrunspersplit=3) 191 csall = CustomSplitter([([0,3,4],[5,9],[2])], 192 nrunspersplit=3) 193 # lets craft simpler dataset 194 #ds = Dataset(samples=N.arange(12), labels=[1]*6+[2]*6, chunks=1) 195 splits = list(cs(self.data)) 196 splitsall = list(csall(self.data)) 197 198 self.failUnless(len(splits) == 3) 199 ul = self.data.uniquelabels 200 201 self.failUnless(((N.array(splitsall[0][0].samplesperlabel.values()) 202 *[0.3, 0.6, 1.0, 0.5]).round().astype(int) == 203 N.array(splits[0][0].samplesperlabel.values())).all()) 204 205 self.failUnless(((N.array(splitsall[0][1].samplesperlabel.values())*0.5 206 ).round().astype(int) == 207 N.array(splits[0][1].samplesperlabel.values())).all()) 208 209 self.failUnless((N.array(splitsall[0][2].samplesperlabel.values()) == 210 N.array(splits[0][2].samplesperlabel.values())).all())
211 212
213 - def testNoneSplitter(self):
214 nos = NoneSplitter() 215 splits = [ (train, test) for (train, test) in nos(self.data) ] 216 self.failUnless(len(splits) == 1) 217 self.failUnless(splits[0][0] == None) 218 self.failUnless(splits[0][1].nsamples == 100) 219 220 nos = NoneSplitter(mode='first') 221 splits = [ (train, test) for (train, test) in nos(self.data) ] 222 self.failUnless(len(splits) == 1) 223 self.failUnless(splits[0][1] == None) 224 self.failUnless(splits[0][0].nsamples == 100) 225 226 227 # test sampling tools 228 # specified value 229 nos = NoneSplitter(nrunspersplit=3, 230 nperlabel=10) 231 splits = [ (train, test) for (train, test) in nos(self.data) ] 232 233 self.failUnless(len(splits) == 3) 234 for split in splits: 235 self.failUnless(split[0] == None) 236 self.failUnless(split[1].nsamples == 40) 237 self.failUnless(split[1].samplesperlabel.values() == [10,10,10,10]) 238 239 # auto-determined 240 nos = NoneSplitter(nrunspersplit=3, 241 nperlabel='equal') 242 splits = [ (train, test) for (train, test) in nos(self.data) ] 243 244 self.failUnless(len(splits) == 3) 245 for split in splits: 246 self.failUnless(split[0] == None) 247 self.failUnless(split[1].nsamples == 100) 248 self.failUnless(split[1].samplesperlabel.values() == [25,25,25,25])
249 250
251 - def testLabelSplitter(self):
252 oes = OddEvenSplitter(attr='labels') 253 254 splits = [ (first, second) for (first, second) in oes(self.data) ] 255 256 self.failUnless((splits[0][0].uniquelabels == [0,2]).all()) 257 self.failUnless((splits[0][1].uniquelabels == [1,3]).all()) 258 self.failUnless((splits[1][0].uniquelabels == [1,3]).all()) 259 self.failUnless((splits[1][1].uniquelabels == [0,2]).all())
260 261
262 - def testCountedSplitting(self):
263 # count > #chunks, should result in 10 splits 264 nchunks = len(self.data.uniquechunks) 265 for strategy in NFoldSplitter._STRATEGIES: 266 for count, target in [ (nchunks*2, nchunks), 267 (nchunks, nchunks), 268 (nchunks-1, nchunks-1), 269 (3, 3), 270 (0, 0), 271 (1, 1) 272 ]: 273 nfs = NFoldSplitter(cvtype=1, count=count, strategy=strategy) 274 splits = [ (train, test) for (train,test) in nfs(self.data) ] 275 self.failUnless(len(splits) == target) 276 chosenchunks = [int(s[1].uniquechunks) for s in splits] 277 if strategy == 'first': 278 self.failUnlessEqual(chosenchunks, range(target)) 279 elif strategy == 'equidistant': 280 if target == 3: 281 self.failUnlessEqual(chosenchunks, [0, 3, 7]) 282 elif strategy == 'random': 283 # none is selected twice 284 self.failUnless(len(set(chosenchunks)) == len(chosenchunks)) 285 self.failUnless(target == len(chosenchunks)) 286 else: 287 raise RuntimeError, "Add unittest for strategy %s" \ 288 % strategy
289 290
291 - def testDiscardedBoundaries(self):
292 splitters = [NFoldSplitter(), 293 NFoldSplitter(discard_boundary=(0,1)), # discard testing 294 NFoldSplitter(discard_boundary=(1,0)), # discard training 295 NFoldSplitter(discard_boundary=(2,0)), # discard 2 from training 296 NFoldSplitter(discard_boundary=1), # discard from both 297 OddEvenSplitter(discard_boundary=(1,0)), 298 OddEvenSplitter(discard_boundary=(0,1)), 299 HalfSplitter(discard_boundary=(1,0)), 300 ] 301 302 split_sets = [list(s(self.data)) for s in splitters] 303 counts = [[(len(s[0].chunks), len(s[1].chunks)) for s in split_set] 304 for split_set in split_sets] 305 306 nodiscard_tr = [c[0] for c in counts[0]] 307 nodiscard_te = [c[1] for c in counts[0]] 308 309 # Discarding in testing: 310 self.failUnless(nodiscard_tr == [c[0] for c in counts[1]]) 311 self.failUnless(nodiscard_te[1:-1] == [c[1] + 2 for c in counts[1][1:-1]]) 312 # at the beginning/end chunks, just a single element 313 self.failUnless(nodiscard_te[0] == counts[1][0][1] + 1) 314 self.failUnless(nodiscard_te[-1] == counts[1][-1][1] + 1) 315 316 # Discarding in training 317 for d in [1,2]: 318 self.failUnless(nodiscard_te == [c[1] for c in counts[1+d]]) 319 self.failUnless(nodiscard_tr[0] == counts[1+d][0][0] + d) 320 self.failUnless(nodiscard_tr[-1] == counts[1+d][-1][0] + d) 321 self.failUnless(nodiscard_tr[1:-1] == [c[0] + d*2 322 for c in counts[1+d][1:-1]]) 323 324 # Discarding in both -- should be eq min from counts[1] and [2] 325 counts_min = [(min(c1[0], c2[0]), min(c1[1], c2[1])) 326 for c1,c2 in zip(counts[1], counts[2])] 327 self.failUnless(counts_min == counts[4])
328 329 # TODO: test all those odd/even etc splitters... YOH: did 330 # visually... looks ok;) 331 #for count in counts[5:]: 332 # print count 333 334
335 -def suite():
336 return unittest.makeSuite(SplitterTests)
337 338 339 if __name__ == '__main__': 340 import runner 341