Работа с несбалансированными выборками

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

Проблемы несбалансированных выборок

  1. Смещение к основному классу Стандартный алгоритм обучения нейронной сети минимизирует ошибку на всей выборке. Если один класс преобладает, сеть будет «учиться» распознавать именно его, игнорируя редкие классы.

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

  3. Проблемы метрик качества Метрики, такие как точность (accuracy), становятся малоинформативными. Модель, предсказывающая всегда доминирующий класс, может иметь высокую точность, но при этом полностью игнорировать редкие классы. Более подходящие метрики: precision, recall, F1-score, матрица ошибок.

Подходы к работе с несбалансированными выборками в Brain.js

1. Взвешивание классов

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

Пример:

const brain = require('brain.js');

const net = new brain.NeuralNetwork({ hiddenLayers: [10, 10] });

const trainingData = [
  { input: { cat: 1 }, output: { yes: 1 } },
  { input: { dog: 1 }, output: { no: 1 } },
  { input: { rareAnimal: 1 }, output: { yes: 1 } },
];

// Увеличиваем вес редкого класса вручную
const weightedData = trainingData.map(item => {
  if (item.input.rareAnimal) {
    return { ...item, weight: 5 };
  }
  return { ...item, weight: 1 };
});

net.train(weightedData, {
  iterations: 20000,
  log: true,
  logPeriod: 1000,
});

Ключевой момент: Brain.js позволяет использовать свойство weight для увеличения влияния редких примеров на обучение.

2. Искусственное расширение выборки (Oversampling)

Если класс редкий, можно создавать дополнительные синтетические примеры. На практике это делается с дублированием редких примеров или генерацией новых комбинаций признаков.

Пример дублирования:

let augmentedData = [...trainingData];
const rareExamples = trainingData.filter(d => d.input.rareAnimal);

for (let i = 0; i < 5; i++) {
  augmentedData = augmentedData.concat(rareExamples);
}

Такой подход помогает сети увидеть больше примеров редкого класса, снижая смещение.

3. Подвыборка (Undersampling) доминирующего класса

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

const dominantClass = trainingData.filter(d => d.input.cat);
const rareClass = trainingData.filter(d => d.input.rareAnimal);

const balancedData = dominantClass.slice(0, rareClass.length * 2).concat(rareClass);

Метод уменьшает перекос, но требует достаточного количества данных для обучения.

4. Генерация дополнительных признаков

В некоторых случаях несбалансированность можно компенсировать увеличением информативности входных данных. Например, кодирование категориальных признаков через one-hot или добавление новых признаков для редких классов.

5. Настройка параметров сети
  • Количество скрытых слоев и нейронов: Для редких классов стоит выбирать более глубокие сети с большим количеством нейронов, чтобы уловить сложные закономерности.
  • Функция активации: Sigmoid может плохо справляться с редкими событиями; в Brain.js можно экспериментировать с ReLU для скрытых слоев.
  • Скорость обучения (learningRate): При несбалансированных данных рекомендуется использовать меньший learningRate, чтобы редкие примеры оказывали более устойчивое влияние.

Пример комплексного подхода

const net = new brain.NeuralNetwork({ hiddenLayers: [20, 10] });

// Взвешивание и увеличение редких классов
const balancedTrainingData = trainingData.flatMap(item => {
  if (item.input.rareAnimal) {
    return Array(5).fill({ ...item, weight: 2 });
  }
  return { ...item, weight: 1 };
});

net.train(balancedTrainingData, {
  iterations: 25000,
  learningRate: 0.01,
  log: true,
  logPeriod: 2000,
});

Такой подход сочетает взвешивание, oversampling и тонкую настройку сети, что позволяет добиться более справедливого распознавания редких классов.

Метрики для оценки работы с несбалансированными данными

  • Матрица ошибок (Confusion Matrix) – наглядно показывает, какие классы сеть путает.
  • Precision и Recall – позволяют оценить точность и полноту предсказаний редких классов.
  • F1-score – гармоническое среднее precision и recall, лучше отражает качество модели при несбалансированных данных.

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


Хотите, я могу подготовить раздел с конкретными стратегиями генерации синтетических данных для Brain.js, чтобы редкие классы всегда присутствовали в достаточном объеме и сеть не игнорировала их?