| 1 | def cat_dataset(batch_size, iters, dataset_size, imgs_per_gpu, sample_weight): |
| 2 | print(f"====> dataset_size:\t{dataset_size}") |
| 3 | print(f"====> sample_weight:\t{sample_weight}") |
| 4 | print(f"====> iters:\t{iters}") |
| 5 | print(f"initial batch_size:\t{batch_size}") |
| 6 | total = sum(dataset_size.values()) |
| 7 | t_wise_batchsize = {k: v / total * batch_size for k, v in dataset_size.items()} |
| 8 | t_wise_gpuse = {k: v / imgs_per_gpu[k] for k, v in t_wise_batchsize.items()} |
| 9 | |
| 10 | print(f"-----------------------rounded-----------------------") |
| 11 | gpus = 1 |
| 12 | while gpus % 8 != 0: |
| 13 | t_wise_batchsize = {k: round(v / total * batch_size / imgs_per_gpu[k]) * imgs_per_gpu[k] for k, v in |
| 14 | dataset_size.items()} |
| 15 | t_wise_gpuse = {k: v // imgs_per_gpu[k] for k, v in t_wise_batchsize.items()} |
| 16 | t_wise_epochs = {k: iters * t_wise_batchsize[k] / v for k, v in dataset_size.items()} |
| 17 | batch_size += 1 |
| 18 | gpus = sum(t_wise_gpuse.values()) |
| 19 | loss_weights = {k: v * t_wise_batchsize[k] for k, v in sample_weight.items()} |
| 20 | print(f"rounded batch_size:\t{sum(t_wise_batchsize.values())}") |
| 21 | print(f"task-wise batchsize:\t{t_wise_batchsize}") |
| 22 | print(f"task-wise epoch:\t{t_wise_epochs}\t avg apoch: {sum(t_wise_epochs.values()) / len(t_wise_epochs)}") |
| 23 | print(f"loss weights:\t{loss_weights}") |
| 24 | |
| 25 | print(f"task-wise gpus:\t{t_wise_gpuse}\ngpus:\t{gpus} ({gpus // 8} nodes)") |
| 26 | |
| 27 | |
| 28 | # example usage |