LSTM (Long Short-Term Memory) — разновидность рекуррентных нейронных сетей (RNN), предназначенная для обработки последовательных данных, таких как текст. Основное отличие LSTM от обычных RNN заключается в механизме ячейки памяти, которая позволяет сети запоминать важную информацию на длительные промежутки времени, избегая проблемы исчезающего градиента.
Ядро LSTM состоит из трёх основных гейтов:
Эта архитектура особенно эффективна для задач классификации текста, где контекст из предыдущих слов влияет на понимание последующих.
Перед подачей текста в LSTM необходимо провести токенизацию и векторизацию.
Токенизация: текст разбивается на токены (слова или символы). В Keras.js обычно используется токенизация на стороне Python с сохранением словаря, который затем импортируется в браузер.
Векторизация: каждый токен преобразуется в
числовой индекс или вектор. Для индексов используется
one-hot encoding или embedding layer, который
создаёт плотные векторы фиксированной размерности.
Пример структуры данных для модели:
{
"x": [[12, 4, 56, 23], [34, 2, 78, 5], ...],
"y": [[1, 0], [0, 1], ...] // категориальные метки
}
Для классификации текста LSTM-модель обычно строится следующим образом:
softmax для многоклассовой классификации или
sigmoid для бинарной.Пример конфигурации модели в Keras (Python), которую можно экспортировать для Keras.js:
from keras.models import Sequential
from keras.layers import Embedding, LSTM, Dense
model = Sequential()
model.add(Embedding(input_dim=10000, output_dim=128))
model.add(LSTM(128))
model.add(Dense(2, activation='softmax'))
model.compile(loss='categorical_crossentropy', optimizer='adam', metrics=['accuracy'])
Модель сохраняется в формате model.json с весами в
model_weights.buf для последующего использования в браузере
через Keras.js.
Keras.js позволяет запускать модели в браузере, используя WebGL для ускорения вычислений.
Загрузка модели и весов:
const KerasJS = require('keras-js');
const model = new KerasJS.Model({
filepaths: {
model: 'model.json',
weights: 'model_weights.buf'
},
gpu: true
});
await model.ready();
Подготовка входных данных:
// Пример: последовательность токенов длиной 50
const inputData = new Float32Array(50);
for (let i = 0; i < 50; i++) inputData[i] = tokens[i];
const input = { input_1: inputData };
Запуск предсказания:
const outputData = await model.predict(input);
console.log('Предсказанная вероятность классов:', outputData);
padding с нулями в начале или конце.dropout и recurrent_dropout предотвращает
переобучение.function padSequence(sequence, maxLength) {
const padded = new Array(maxLength).fill(0);
for (let i = 0; i < Math.min(sequence.length, maxLength); i++) {
padded[i] = sequence[i];
}
return padded;
}
const paddedInput = padSequence(tokens, 50);
Для анализа модели в браузере удобно отображать вероятности классов в виде гистограмм или цветовых шкал. Keras.js возвращает Float32Array, который легко интегрировать с библиотеками визуализации, такими как Chart.js или D3.js.