| 308 | } |
| 309 | |
| 310 | float train_network(network *net, data d) |
| 311 | { |
| 312 | assert(d.X.rows % net->batch == 0); |
| 313 | int batch = net->batch; |
| 314 | int n = d.X.rows / batch; |
| 315 | |
| 316 | int i; |
| 317 | float sum = 0; |
| 318 | for(i = 0; i < n; ++i){ |
| 319 | get_next_batch(d, batch, i*batch, net->input, net->truth); |
| 320 | float err = train_network_datum(net); |
| 321 | sum += err; |
| 322 | } |
| 323 | return (float)sum/(n*batch); |
| 324 | } |
| 325 | |
| 326 | void set_temp_network(network *net, float t) |
| 327 | { |
no test coverage detected