MCPcopy Create free account
hub / github.com/Elsaam2y/DINet_optimized / DataParallelWithCallback

Class DataParallelWithCallback

sync_batchnorm/replicate.py:50–67  ·  view source on GitHub ↗

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

Source from the content-addressed store, hash-verified

48
49
50class 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
70def patch_replication_callback(data_parallel):

Callers 1

convert_modelFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected