| 145 | |
| 146 | # Data partition according to the rank |
| 147 | def partition(global_rank, world_size, train_x, train_y, val_x, val_y): |
| 148 | # Partition training data |
| 149 | data_per_rank = train_x.shape[0] // world_size |
| 150 | idx_start = global_rank * data_per_rank |
| 151 | idx_end = (global_rank + 1) * data_per_rank |
| 152 | train_x = train_x[idx_start:idx_end] |
| 153 | train_y = train_y[idx_start:idx_end] |
| 154 | |
| 155 | # Partition evaluation data |
| 156 | data_per_rank = val_x.shape[0] // world_size |
| 157 | idx_start = global_rank * data_per_rank |
| 158 | idx_end = (global_rank + 1) * data_per_rank |
| 159 | val_x = val_x[idx_start:idx_end] |
| 160 | val_y = val_y[idx_start:idx_end] |
| 161 | return train_x, train_y, val_x, val_y |
| 162 | |
| 163 | |
| 164 | # Function to all reduce NUMPY accuracy and loss from multiple devices |