В TensorFlow.js обучение нейронной сети обычно осуществляется с помощью встроенного оптимизатора, который автоматически обновляет веса модели на каждом шаге обратного распространения ошибки. Однако существует необходимость полного контроля над процессом обновления весов, например, для реализации кастомных алгоритмов обучения, экспериментов с динамическими скоростями обучения или нестандартных методов оптимизации.
Обновление весов нейронной сети — это процесс корректировки параметров модели на основе вычисленных градиентов функции потерь. Формально, для веса ( w ) и функции потерь ( L ) шаг обновления с использованием градиентного спуска выглядит так:
[ w := w - ]
где ( ) — коэффициент обучения (learning rate), а ( ) — градиент функции потерь по весу ( w ).
В TensorFlow.js градиенты вычисляются с помощью функции
tf.grad или tf.variableGrads, что позволяет
использовать их для ручного обновления параметров модели.
Для ручного управления весами все параметры модели должны быть
объявлены как tf.Variable. Пример:
const w1 = tf.variable(tf.randomNormal([inputSize, hiddenSize]));
const b1 = tf.variable(tf.zeros([hiddenSize]));
const w2 = tf.variable(tf.randomNormal([hiddenSize, outputSize]));
const b2 = tf.variable(tf.zeros([outputSize]));
Использование tf.variable позволяет изменять значения
тензоров после их создания, что необходимо для прямого применения
градиентного спуска вручную.
Для вычисления градиентов используется функция
tf.variableGrads, которая принимает функцию потерь и
возвращает объект с градиентами для всех переменных:
const loss = (pred, label) => pred.sub(label).square().mean();
const computeGradients = (inputs, labels) => {
return tf.variableGrads(() => {
const hidden = inputs.matMul(w1).add(b1).relu();
const output = hidden.matMul(w2).add(b2);
return loss(output, labels);
});
};
В объекте, возвращаемом tf.variableGrads, ключи — это
переменные, а значения — соответствующие градиенты.
После вычисления градиентов веса можно обновлять самостоятельно. Например, используя классический градиентный спуск:
const learningRate = 0.01;
const applyGradients = (grads) => {
w1.assign(w1.sub(grads[w1].mul(learningRate)));
b1.assign(b1.sub(grads[b1].mul(learningRate)));
w2.assign(w2.sub(grads[w2].mul(learningRate)));
b2.assign(b2.sub(grads[b2].mul(learningRate)));
};
assign заменяет текущее значение переменной новым
тензором. Выражение w1.sub(grads[w1].mul(learningRate))
создаёт обновлённый тензор для веса с учётом шага обучения.
tf.tidy для управления памятьюПри ручном обновлении весов важно контролировать использование
памяти, так как TensorFlow.js не освобождает автоматически промежуточные
тензоры. Для этого все вычисления можно помещать в
tf.tidy:
tf.tidy(() => {
const grads = computeGradients(inputs, labels);
applyGradients(grads.grads);
});
tf.tidy гарантирует удаление всех промежуточных
тензоров, которые больше не нужны, предотвращая утечки памяти в
браузере.
Ручное управление весами позволяет реализовать любой алгоритм оптимизации. Например, стохастический градиентный спуск с импульсом:
let vW1 = tf.zerosLike(w1);
let vB1 = tf.zerosLike(b1);
let vW2 = tf.zerosLike(w2);
let vB2 = tf.zerosLike(b2);
const momentum = 0.9;
const applyMomentum = (grads) => {
vW1 = vW1.mul(momentum).sub(grads[w1].mul(learningRate));
vB1 = vB1.mul(momentum).sub(grads[b1].mul(learningRate));
vW2 = vW2.mul(momentum).sub(grads[w2].mul(learningRate));
vB2 = vB2.mul(momentum).sub(grads[b2].mul(learningRate));
w1.assign(w1.add(vW1));
b1.assign(b1.add(vB1));
w2.assign(w2.add(vW2));
b2.assign(b2.add(vB2));
};
Использование импульса помогает ускорить сходимость и сгладить колебания градиентов.
Полный цикл обучения с ручным обновлением весов выглядит следующим образом:
for (let epoch = 0; epoch < epochs; epoch++) {
tf.tidy(() => {
const grads = computeGradients(inputs, labels);
applyGradients(grads.grads);
});
if (epoch % 10 === 0) {
const currentLoss = loss(forward(inputs), labels).dataSync()[0];
console.log(`Epoch ${epoch}: Loss = ${currentLoss}`);
}
}
forward(inputs) — функция прямого прохода, аналогичная
вычислению предсказания в computeGradients. Этот цикл
позволяет полностью контролировать процесс обновления весов на каждом
шаге.
Использование этих подходов в TensorFlow.js открывает путь к глубокому пониманию работы нейронных сетей и расширяет возможности экспериментов с обучением моделей на стороне клиента.