MCPcopy Create free account
hub / github.com/apache/singa / Train

Function Train

examples/cpp/cifar10/cnn-parallel.cc:142–247  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

140}
141
142void Train(float lr, int num_epoch, string data_dir) {
143 Cifar10 data(data_dir);
144 Tensor train_x, train_y, test_x, test_y;
145 Tensor train_x_1, train_x_2, train_y_1, train_y_2;
146 {
147 auto train = data.ReadTrainData();
148 size_t nsamples = train.first.shape(0);
149 auto mtrain =
150 Reshape(train.first, Shape{nsamples, train.first.Size() / nsamples});
151 const Tensor &mean = Average(mtrain, 0);
152 SubRow(mean, &mtrain);
153 train_x = Reshape(mtrain, train.first.shape());
154 train_y = train.second;
155
156 LOG(INFO) << "Slicing training data...";
157 train_x_1 = Tensor(Shape{nsamples / 2, train.first.shape(1),
158 train.first.shape(2), train.first.shape(3)});
159 LOG(INFO) << "Copying first data slice...";
160 CopyDataToFrom(&train_x_1, train_x, train_x.Size() / 2);
161 train_x_2 = Tensor(Shape{nsamples / 2, train.first.shape(1),
162 train.first.shape(2), train.first.shape(3)});
163 LOG(INFO) << "Copying second data slice...";
164 CopyDataToFrom(&train_x_2, train_x, train_x.Size() / 2, 0,
165 train_x.Size() / 2);
166 train_y_1 = Tensor(Shape{nsamples / 2});
167 train_y_1.AsType(kInt);
168 LOG(INFO) << "Copying first label slice...";
169 CopyDataToFrom(&train_y_1, train_y, train_y.Size() / 2);
170 train_y_2 = Tensor(Shape{nsamples / 2});
171 train_y_2.AsType(kInt);
172 LOG(INFO) << "Copying second label slice...";
173 CopyDataToFrom(&train_y_2, train_y, train_y.Size() / 2, 0,
174 train_y.Size() / 2);
175
176 auto test = data.ReadTestData();
177 nsamples = test.first.shape(0);
178 auto mtest =
179 Reshape(test.first, Shape{nsamples, test.first.Size() / nsamples});
180 SubRow(mean, &mtest);
181 test_x = Reshape(mtest, test.first.shape());
182 test_y = test.second;
183 }
184
185 CHECK_EQ(train_x.shape(0), train_y.shape(0));
186 CHECK_EQ(test_x.shape(0), test_y.shape(0));
187 LOG(INFO) << "Total Training samples = " << train_y.shape(0)
188 << ", Total Test samples = " << test_y.shape(0);
189 CHECK_EQ(train_x_1.shape(0), train_y_1.shape(0));
190 LOG(INFO) << "On net 1, Training samples = " << train_y_1.shape(0)
191 << ", Test samples = " << test_y.shape(0);
192 CHECK_EQ(train_x_2.shape(0), train_y_2.shape(0));
193 LOG(INFO) << "On net 2, Training samples = " << train_y_2.shape(0);
194
195 auto net_1 = CreateNet();
196 auto net_2 = CreateNet();
197
198 SGD sgd;
199 OptimizerConf opt_conf;

Callers 1

mainFunction · 0.70

Calls 15

AverageFunction · 0.85
SubRowFunction · 0.85
CopyDataToFromFunction · 0.85
ReadTrainDataMethod · 0.80
shapeMethod · 0.80
ReadTestDataMethod · 0.80
CompileMethod · 0.80
TrainThreadMethod · 0.80
CreateNetFunction · 0.70
ReshapeClass · 0.50
TensorClass · 0.50

Tested by

no test coverage detected