updated args

This commit is contained in:
William Falcon
2019-06-26 17:53:05 -04:00
parent 0b1e22ac51
commit 4a3c9de857
+1 -1
View File
@@ -399,7 +399,7 @@ class Trainer(TrainerIO):
# when DP, we need to aggregate the scalars we received as outputs
# use mean as the reduce function
if self.data_parallel:
output = reduce_distributed_output(output, len(self.gpus))
output = reduce_distributed_output(output, len(self.data_parallel_device_ids))
model_specific_tqdm_metrics_dic = output['tqdm_metrics']
loss = output['loss']