Samples a fixed number of samples from the dataset, deterministically. Arguments: data_source, sample_size, seed (optional)
| 44 | |
| 45 | |
| 46 | class FixedRandomSubsetSampler(FixedSubsetSampler): |
| 47 | """Samples a fixed number of samples from the dataset, deterministically. |
| 48 | Arguments: |
| 49 | data_source, |
| 50 | sample_size, |
| 51 | seed (optional) |
| 52 | """ |
| 53 | def __init__(self, data_source, start=None, end=None, seed=1): |
| 54 | rng = random.Random(seed) |
| 55 | shuffled = list(range(len(data_source))) |
| 56 | rng.shuffle(shuffled) |
| 57 | self.data_source = data_source |
| 58 | super(FixedRandomSubsetSampler, self).__init__(shuffled[start:end]) |
| 59 | |
| 60 | def class_subset(self, class_filter): |
| 61 | ''' |
| 62 | Returns only the subset matching the given rule. |
| 63 | ''' |
| 64 | if isinstance(class_filter, int): |
| 65 | rule = lambda d: d[1] == class_filter |
| 66 | else: |
| 67 | rule = class_filter |
| 68 | return self.subset([i for i, j in enumerate(self.samples) |
| 69 | if rule(self.data_source[j])]) |
| 70 | |
| 71 | def coordinate_sample(shape, sample_size, seeds, grid=13, seed=1, flat=False): |
| 72 | ''' |