MCPcopy Create free account
hub / github.com/FidoProject/Fido / trainOnDataPoint

Method trainOnDataPoint

src/SGDTrainer.cpp:67–93  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

65
66
67double SGDTrainer::trainOnDataPoint(net::NeuralNet *network, const std::vector<double> &input, const std::vector<double> &correctOutput) {
68 std::vector< std::vector< std::vector<double> > > weights = network->getWeights3D();
69 gradients.push_back(network->getGradients(input, correctOutput));
70
71 // Update weights
72 std::vector< std::vector< std::vector<double> > > newWeightChanges(weights);
73 for(unsigned int layerIndex = 0; layerIndex < weights.size(); layerIndex++) {
74 for(unsigned int neuronIndex = 0; neuronIndex < weights[layerIndex].size(); neuronIndex++) {
75 for(unsigned int weightIndex = 0; weightIndex < weights[layerIndex][neuronIndex].size(); weightIndex++) {
76 double deltaWeight = getChangeInWeight(weights[layerIndex][neuronIndex][weightIndex], layerIndex, neuronIndex, weightIndex);
77 weights[layerIndex][neuronIndex][weightIndex] += deltaWeight;
78 newWeightChanges[layerIndex][neuronIndex][weightIndex] = deltaWeight;
79 }
80 }
81 }
82 weightChanges.push_back(newWeightChanges);
83
84 network->setWeights3D(weights);
85
86 double networkError = 0;
87 std::vector<double> output = network->getOutput(input);
88 for(unsigned int outputIndex = 0; outputIndex < output.size(); outputIndex++) {
89 networkError += pow(correctOutput[outputIndex] - output[outputIndex], 2);
90 }
91
92 return networkError;
93}
94
95
96void SGDTrainer::resetNetworkVectors(net::NeuralNet *network) {

Callers

nothing calls this directly

Calls 5

getWeights3DMethod · 0.80
setWeights3DMethod · 0.80
getGradientsMethod · 0.45
sizeMethod · 0.45
getOutputMethod · 0.45

Tested by

no test coverage detected