1+ struct train_stage_time_report {
2+ Timer total;
3+ Timer loader;
4+ Timer model;
5+ Timer reduce;
6+
7+ train_stage_time_report (): total(true ){}
8+
9+ std::string report (double FLOPS ){
10+ double Gflops = FLOPS /1e9 /model.time ();
11+ std::ostringstream os;
12+ double other = total.time () - loader.time () - model.time () - reduce.time ();
13+ os << " loader: " << loader.time () << " s, "
14+ << " model: " << model.time () << " s (" << Gflops << " Gflops), "
15+ << " reduction: " << reduce.time () << " s, "
16+ << " other: " << other << " s | "
17+ << " total: " << total.time ();
18+ return os.str ();
19+ }
20+ };
21+
22+
23+
124template <typename DataLoader, typename LossWrappedModelType, typename Optimizer>
225std::pair<
326std::vector<typename LossWrappedModelType::FloatType>,
@@ -50,8 +73,8 @@ train(LossWrappedModelType &loss_func, const DataLoader &train_data, DataLoader
5073 // ////////// train epoch ///////////////
5174 optimizer.epochStart (epoch, do_print);
5275 std::random_shuffle ( didx_train.begin (), didx_train.end (), [&](const int l){ return dist (gen); } ); // shuffle training data indices
53- FloatType lmax_train=std::numeric_limits<FloatType>::lowest (), lmin_train = std::numeric_limits<FloatType>::max (), lavg_train = 0 .;
54- auto ts= now () ;
76+ FloatType lmax_train=std::numeric_limits<FloatType>::lowest (), lmin_train = std::numeric_limits<FloatType>::max (), lavg_train = 0 .;
77+ train_stage_time_report t_train ;
5578
5679 for (int block=0 ;block<nblocks_ddp_train;block++){
5780 int ddp_blocksize_actual = std::min (nbatch_train - block*ddp_blocksize, ddp_blocksize);
@@ -63,14 +86,21 @@ train(LossWrappedModelType &loss_func, const DataLoader &train_data, DataLoader
6386 int bidx = block*ddp_blocksize + me; // which batch are we doing?
6487
6588 // Get the batch
66- auto bxy = train_data.batch (didx_train.data () + bidx*batch_size, batch_size);
89+ TIME (t_train.loader ,
90+ auto bxy = train_data.batch (didx_train.data () + bidx*batch_size, batch_size);
91+ );
6792
93+ TIME (t_train.model ,
6894 loss = loss_func.loss (bxy.x , bxy.y , DerivYes);
6995 deriv = loss_func.deriv ();
96+ );
7097 }
98+
99+ TIME (t_train.reduce ,
71100 ddpAverage (&loss,1 ,false ); // no need to bcast the loss to the pipeline ranks
72101 ddpAverage (deriv,true ); // share the deriv over all pipeline ranks
73-
102+ )
103+
74104 // if(do_print) std::cout << epoch << "-" << block << " : "<< loss << std::endl;
75105 lmax_train = std::max (lmax_train,loss);
76106 lmin_train = std::min (lmin_train,loss);
@@ -84,45 +114,59 @@ train(LossWrappedModelType &loss_func, const DataLoader &train_data, DataLoader
84114 losses_train[block+nblocks_ddp_train*epoch] = loss;
85115 }
86116 lavg_train /= nblocks_ddp_train;
87- double train_time = since (ts);
88- double train_Tflops = nbatch_train * double (loss_func.FLOPS (0 ) + loss_func.FLOPS (1 ) + 2 *loss_func.nparams ())/1.0e12 / train_time;
117+
118+ t_train.total .pause ();
119+ double train_FLOPS = nbatch_train * double (loss_func.FLOPS (0 ) + loss_func.FLOPS (1 ));
89120
90121 // ////////// end train epoch ///////////////
91122
92123 // ////////// validate epoch ///////////////
93124 if (valid_data){
94125 FloatType lmax_valid=std::numeric_limits<FloatType>::lowest (), lmin_valid = std::numeric_limits<FloatType>::max (), lavg_valid = 0 .;
95- ts= now () ;
126+ train_stage_time_report t_valid ;
96127
97128 for (int block=0 ;block<nblocks_ddp_valid;block++){
98129 int ddp_blocksize_actual = std::min (nbatch_valid - block*ddp_blocksize, ddp_blocksize);
99130
100131 FloatType loss = 0 ;
101132 if (me < ddp_blocksize_actual){
102- int bidx = block*ddp_blocksize + me;
133+ int bidx = block*ddp_blocksize + me;
134+ TIME (t_valid.loader ,
103135 auto bxy = valid_data->batch (didx_valid.data () + bidx*batch_size, batch_size); // no need to shuffle
136+ );
137+ TIME (t_valid.model ,
104138 loss = loss_func.loss (bxy.x , bxy.y , DerivNo);
139+ );
105140 }
106-
141+
142+ TIME (t_valid.reduce ,
107143 ddpAverage (&loss,1 ,false );
108-
144+ );
145+
109146 lmax_valid = std::max (lmax_valid,loss);
110147 lmin_valid = std::min (lmin_valid,loss);
111148 lavg_valid += loss;
112149
113150 losses_valid[block+nblocks_ddp_valid*epoch] = loss;
114- }
151+ }
115152 lavg_valid /= nblocks_ddp_valid;
116- double valid_time = since (ts);
117- double valid_Tflops = nbatch_valid * double (loss_func.FLOPS (0 ))/1.0e12 / valid_time;
153+ t_valid.total .pause ();
154+
155+ double valid_FLOPS = nbatch_valid * double (loss_func.FLOPS (0 ));
118156
119157 // ////////// end validate epoch ///////////////
120158
121159 if (do_print) std::cout << " Epoch : " << epoch << std::endl
122- << " training time : " << train_time <<" s (" << train_Tflops << " Tflops) loss min: " << lmin_train << " avg: " << lavg_train << " max: " << lmax_train << std::endl
123- << " validation time : " << valid_time <<" s (" << valid_Tflops << " Tflops) loss min: " << lmin_valid << " avg: " << lavg_valid << " max: " << lmax_valid << std::endl;
124- }else { // if not validating, just print info on the training losses
125- if (do_print) std::cout << " Epoch : " << epoch << " time : " << train_time <<" s (" << train_Tflops << " Tflops) loss min: " << lmin_train << " avg: " << lavg_train << " max: " << lmax_train << std::endl;
160+ << " training loss min: " << lmin_train << " avg: " << lavg_train << " max: " << lmax_train << std::endl
161+ << " validation loss min: " << lmin_valid << " avg: " << lavg_valid << " max: " << lmax_valid << std::endl
162+ << " training timings: " << t_train.report (train_FLOPS) << std::endl
163+ << " validation timings: " << t_valid.report (valid_FLOPS) << std::endl;
164+
165+
166+ }else { // if not validating, just print info on the training losses
167+ if (do_print) std::cout << " Epoch : " << epoch << std::endl
168+ << " loss min: " << lmin_train << " avg: " << lavg_train << " max: " << lmax_train << std::endl
169+ << " timings: " << t_train.report (train_FLOPS) << std::endl;
126170 }
127171 }// epoch
128172
0 commit comments