mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
Fix build error when USE_NCCL is defined. (#4334)
This commit is contained in:
parent
0d9db2b28d
commit
a6d10376df
1 changed files with 4 additions and 4 deletions
|
|
@ -412,7 +412,7 @@ TEST_F(OptimizerGraphBuilderTest, ZeRO_NoGradientAccumulation_NoMixedPrecision)
|
|||
OptimizerGraphConfig config;
|
||||
config.data_parallel_group_size = 4;
|
||||
config.use_nccl = true;
|
||||
config.deepspeed_config = ZeROConfig{0};
|
||||
config.deepspeed_zero = ZeROConfig{0};
|
||||
config.gradient_accumulation_steps = 1;
|
||||
config.use_mixed_precision = false;
|
||||
TestZeROOptimizerGraphBuilder(config, graph_);
|
||||
|
|
@ -422,7 +422,7 @@ TEST_F(OptimizerGraphBuilderTest, ZeRO_WithGradientAccumulation_NoMixedPrecision
|
|||
OptimizerGraphConfig config;
|
||||
config.data_parallel_group_size = 4;
|
||||
config.use_nccl = true;
|
||||
config.deepspeed_config = ZeROConfig{0};
|
||||
config.deepspeed_zero = ZeROConfig{0};
|
||||
config.gradient_accumulation_steps = 10;
|
||||
config.use_mixed_precision = false;
|
||||
TestZeROOptimizerGraphBuilder(config, graph_);
|
||||
|
|
@ -432,7 +432,7 @@ TEST_F(OptimizerGraphBuilderTest, ZeRO_NoGradientAccumulation_WithMixedPrecision
|
|||
OptimizerGraphConfig config;
|
||||
config.data_parallel_group_size = 4;
|
||||
config.use_nccl = true;
|
||||
config.deepspeed_config = ZeROConfig{0};
|
||||
config.deepspeed_zero = ZeROConfig{0};
|
||||
config.gradient_accumulation_steps = 1;
|
||||
config.use_mixed_precision = true;
|
||||
config.loss_scale_input_name = k_loss_scaling_factor_name;
|
||||
|
|
@ -443,7 +443,7 @@ TEST_F(OptimizerGraphBuilderTest, ZeRO_WithGradientAccumulation_WithMixedPrecisi
|
|||
OptimizerGraphConfig config;
|
||||
config.data_parallel_group_size = 4;
|
||||
config.use_nccl = true;
|
||||
config.deepspeed_config = ZeROConfig{0};
|
||||
config.deepspeed_zero = ZeROConfig{0};
|
||||
config.gradient_accumulation_steps = 10;
|
||||
config.use_mixed_precision = true;
|
||||
config.loss_scale_input_name = k_loss_scaling_factor_name;
|
||||
|
|
|
|||
Loading…
Reference in a new issue