UNPKG

@tensorflow/tfjs-core

Version:

Hardware-accelerated JavaScript library for machine intelligence

83 lines 2.9 kB
/** * @license * Copyright 2018 Google LLC. All Rights Reserved. * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. * You may obtain a copy of the License at * * http://www.apache.org/licenses/LICENSE-2.0 * * Unless required by applicable law or agreed to in writing, software * distributed under the License is distributed on an "AS IS" BASIS, * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. * See the License for the specific language governing permissions and * limitations under the License. * ============================================================================= */ import { ENGINE } from '../engine'; import { keep, tidy } from '../globals'; import { add } from '../ops/add'; import { mul } from '../ops/mul'; import { scalar } from '../ops/scalar'; import { registerClass } from '../serialization'; import { Optimizer } from './optimizer'; /** @doclink Optimizer */ export class SGDOptimizer extends Optimizer { constructor(learningRate) { super(); this.learningRate = learningRate; this.setLearningRate(learningRate); } applyGradients(variableGradients) { const varNames = Array.isArray(variableGradients) ? variableGradients.map(v => v.name) : Object.keys(variableGradients); varNames.forEach((name, i) => { const gradient = Array.isArray(variableGradients) ? variableGradients[i].tensor : variableGradients[name]; if (gradient == null) { return; } const value = ENGINE.registeredVariables[name]; tidy(() => { const newValue = add(mul(this.c, gradient), value); value.assign(newValue); }); }); this.incrementIterations(); } /** * Sets the learning rate of the optimizer. */ setLearningRate(learningRate) { this.learningRate = learningRate; if (this.c != null) { this.c.dispose(); } this.c = keep(scalar(-learningRate)); } dispose() { this.c.dispose(); } async getWeights() { return [await this.saveIterations()]; } async setWeights(weightValues) { weightValues = await this.extractIterations(weightValues); if (weightValues.length !== 0) { throw new Error('SGD optimizer does not have settable weights.'); } } getConfig() { return { 'learningRate': this.learningRate }; } /** @nocollapse */ static fromConfig(cls, config) { return new cls(config['learningRate']); } } /** @nocollapse */ SGDOptimizer.className = 'SGD'; // Note: Name matters for Python compatibility. registerClass(SGDOptimizer); //# sourceMappingURL=sgd_optimizer.js.map