Data Parallel with a replication callback. An replication callback `__data_parallel_replicate__` of each module will be invoked after being created by original `replicate` function. The callback will be invoked with arguments `__data_parallel_replicate__(ctx, copy_id)` Example
| 48 | |
| 49 | |
| 50 | class DataParallelWithCallback(DataParallel): |
| 51 | """ |
| 52 | Data Parallel with a replication callback. |
| 53 | |
| 54 | An replication callback `__data_parallel_replicate__` of each module will be invoked after being created by |
| 55 | original `replicate` function. |
| 56 | The callback will be invoked with arguments `__data_parallel_replicate__(ctx, copy_id)` |
| 57 | |
| 58 | Examples: |
| 59 | > sync_bn = SynchronizedBatchNorm1d(10, eps=1e-5, affine=False) |
| 60 | > sync_bn = DataParallelWithCallback(sync_bn, device_ids=[0, 1]) |
| 61 | # sync_bn.__data_parallel_replicate__ will be invoked. |
| 62 | """ |
| 63 | |
| 64 | def replicate(self, module, device_ids): |
| 65 | modules = super(DataParallelWithCallback, self).replicate(module, device_ids) |
| 66 | execute_replication_callbacks(modules) |
| 67 | return modules |
| 68 | |
| 69 | |
| 70 | def patch_replication_callback(data_parallel): |