1
2
3
4
5
6
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
20
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
29
30 nfs = NFoldSplitter(cvtype=1)
31
32
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
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
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
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
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
94
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
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
118 s5 = NGroupSplitter(5)
119
120
121 splits = [ (train, test) for (train, test) in s5(self.data) ]
122
123
124 self.failUnless(len(splits) == 5)
125
126
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
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
142
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
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
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
185
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
194
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
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
228
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
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
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
263
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
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
292 splitters = [NFoldSplitter(),
293 NFoldSplitter(discard_boundary=(0,1)),
294 NFoldSplitter(discard_boundary=(1,0)),
295 NFoldSplitter(discard_boundary=(2,0)),
296 NFoldSplitter(discard_boundary=1),
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
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
313 self.failUnless(nodiscard_te[0] == counts[1][0][1] + 1)
314 self.failUnless(nodiscard_te[-1] == counts[1][-1][1] + 1)
315
316
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
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
330
331
332
333
334
337
338
339 if __name__ == '__main__':
340 import runner
341