| 127 | } |
| 128 | |
| 129 | pub fn split(&self, n_samples: usize) -> Result<Vec<(Vec<usize>, Vec<usize>)>, String> { |
| 130 | if n_samples != self.samples_info_sets.len() { |
| 131 | return Err("Dataset length must match samples_info_sets".into()); |
| 132 | } |
| 133 | let n = n_samples; |
| 134 | let mut fold_sizes = vec![n / self.n_splits; self.n_splits]; |
| 135 | for i in 0..(n % self.n_splits) { |
| 136 | fold_sizes[i] += 1; |
| 137 | } |
| 138 | let mut current = 0; |
| 139 | let mut splits = Vec::new(); |
| 140 | for fold_size in fold_sizes { |
| 141 | let start = current; |
| 142 | let stop = current + fold_size; |
| 143 | let test_indices: Vec<usize> = (start..stop).collect(); |
| 144 | let mut train_mask = vec![true; n]; |
| 145 | |
| 146 | // purge overlaps |
| 147 | let test_start = self.samples_info_sets[test_indices[0]].1; |
| 148 | let test_end = self.samples_info_sets[*test_indices.last().unwrap()].1; |
| 149 | for (i, (s, e)) in self.samples_info_sets.iter().enumerate() { |
| 150 | let start_in = *s >= test_start && *s <= test_end; |
| 151 | let end_in = *e >= test_start && *e <= test_end; |
| 152 | let envelop = *s <= test_start && *e >= test_end; |
| 153 | if start_in || end_in || envelop { |
| 154 | train_mask[i] = false; |
| 155 | } |
| 156 | } |
| 157 | |
| 158 | // embargo |
| 159 | let embargo = (self.pct_embargo * n as f64).ceil() as isize; |
| 160 | if embargo > 0 { |
| 161 | let after = (stop as isize + embargo).min(n as isize); |
| 162 | let before = (start as isize - embargo).max(0); |
| 163 | for i in start..(after as usize) { |
| 164 | if i < n { |
| 165 | train_mask[i] = false; |
| 166 | } |
| 167 | } |
| 168 | for i in before as usize..start { |
| 169 | train_mask[i] = false; |
| 170 | } |
| 171 | } |
| 172 | |
| 173 | let train_indices: Vec<usize> = train_mask |
| 174 | .iter() |
| 175 | .enumerate() |
| 176 | .filter_map(|(i, keep)| if *keep { Some(i) } else { None }) |
| 177 | .collect(); |
| 178 | splits.push((train_indices, test_indices)); |
| 179 | current = stop; |
| 180 | } |
| 181 | Ok(splits) |
| 182 | } |
| 183 | } |