Классификация текста

TensorFlow.js — это мощная библиотека для работы с машинным обучением непосредственно в браузере или на Node.js. Одним из ключевых сценариев её применения является классификация текста, когда требуется определить категорию текста: спам/не спам, положительная/отрицательная оценка, тематическая категория и так далее.


Подготовка данных

Для классификации текста необходимо сначала преобразовать текстовые данные в формат, который может обработать нейронная сеть. В TensorFlow.js используются токенизация и векторизация.

Токенизация — это разбиение текста на токены (слова или символы). В простейшем виде:

const sentences = [
  "Я люблю машинное обучение",
  "TensorFlow.js очень удобен"
];

const tokenizer = new tf.layers.textVectorization({outputMode: 'int'});
tokenizer.adapt(sentences);
const sequences = tokenizer.call(sentences);

Векторизация преобразует токены в числовые массивы, пригодные для подачи на вход модели. Один из стандартных подходов — Bag of Words или Word Embeddings (например, tfjs-models/universal-sentence-encoder).

import * as use from '@tensorflow-models/universal-sentence-encoder';

const modelUSE = await use.load();
const embeddings = await modelUSE.embed(sentences);

Эмбеддинги создают плотные векторные представления текста, учитывающие семантическую близость слов.


Построение модели

Для задач классификации текста чаще всего применяются полносвязные нейронные сети (Dense) или рекуррентные сети (RNN, LSTM).

Пример простейшей модели с Dense-слоями:

const model = tf.sequential();

model.add(tf.layers.dense({inputShape: [embeddingSize], units: 128, activation: 'relu'}));
model.add(tf.layers.dropout({rate: 0.5}));
model.add(tf.layers.dense({units: numClasses, activation: 'softmax'}));

model.compile({
  optimizer: tf.train.adam(),
  loss: 'categoricalCrossentropy',
  metrics: ['accuracy']
});

Ключевые моменты:

  • inputShape должен соответствовать размерности входного вектора (например, размер эмбеддинга).
  • Dropout помогает уменьшить переобучение.
  • softmax используется для многоклассовой классификации; для бинарной достаточно sigmoid.

Подготовка меток

Метки категорий преобразуются в one-hot encoding:

const labels = tf.tensor2d([
  [1, 0, 0],
  [0, 1, 0],
  [0, 0, 1]
]);

Для бинарной классификации метки могут быть простыми числами 0 или 1.


Обучение модели

await model.fit(embeddings, labels, {
  epochs: 20,
  batchSize: 32,
  validationSplit: 0.2,
  callbacks: tf.callbacks.earlyStopping({monitor: 'val_loss'})
});

Особенности обучения:

  • validationSplit позволяет следить за точностью на части данных, не использованной для тренировки.
  • earlyStopping предотвращает переобучение, останавливая обучение при отсутствии улучшений на валидации.

Предсказание и интерпретация результатов

После обучения модель готова к предсказанию:

const testSentences = ["Я не люблю ошибки"];
const testEmbeddings = await modelUSE.embed(testSentences);
const predictions = model.predict(testEmbeddings);

predictions.print();
  • Результат predictions — массив вероятностей для каждой категории.
  • Для определения итоговой метки применяется argMax:
const predictedClass = predictions.argMax(-1).dataSync()[0];

Сохранение и загрузка модели

TensorFlow.js позволяет сохранять модель как в браузере, так и на Node.js:

// Сохранение в локальное хранилище браузера
await model.save('localstorage://text-classification-model');

// Загрузка модели
const loadedModel = await tf.loadLayersModel('localstorage://text-classification-model');

Это удобно для повторного использования без необходимости переобучения.


Улучшение качества модели

  1. Очистка текста: удаление пунктуации, приведение к нижнему регистру, токенизация по словам.
  2. Использование предобученных эмбеддингов: Universal Sentence Encoder или Word2Vec повышает семантическое понимание.
  3. Увеличение объема данных: добавление синонимов, перефразирование текстов.
  4. Тонкая настройка модели: подбор числа слоев, нейронов, скорости обучения (learningRate) и функции активации.

Интеграция с веб-приложением

TensorFlow.js позволяет выполнять классификацию в реальном времени в браузере, что исключает необходимость сервера для предсказаний.

Пример использования с элементом формы:

const inputField = document.getElementById('text-input');
inputField.addEventListener('input', async (e) => {
  const text = e.target.value;
  const embedding = await modelUSE.embed([text]);
  const prediction = model.predict(embedding);
  console.log('Предсказание:', prediction.argMax(-1).dataSync()[0]);
});

В результате система способна динамически анализировать текст, классифицируя его по категориям мгновенно.


TensorFlow.js обеспечивает полный цикл разработки текстовой классификации: от подготовки данных и обучения модели до интеграции с веб-интерфейсом. Правильная токенизация, использование эмбеддингов и тщательная настройка сети являются ключевыми факторами успеха в получении точных предсказаний.