Юнит-тестирование моделей с помощью Jest

Для проведения юнит-тестирования моделей Keras.js в JavaScript требуется правильно подготовленное окружение. Основные компоненты:

  1. Node.js версии 16 и выше — обеспечивает выполнение кода и работу тестовых фреймворков.
  2. Keras.js — библиотека для запуска предварительно обученных моделей Keras в браузере или Node.js.
  3. Jest — популярный тестовый фреймворк, обеспечивающий написание и выполнение тестов.

Установка необходимых пакетов через npm:

npm install keras-js jest @tensorflow/tfjs-node
  • keras-js позволяет загружать модели, экспортированные из Keras.
  • @tensorflow/tfjs-node ускоряет выполнение математических операций через нативные оптимизации.
  • jest обеспечивает удобный синтаксис для описания тестов и проверки результатов.

В package.json необходимо добавить скрипт для запуска тестов:

"scripts": {
  "test": "jest"
}

Загрузка и инициализация модели

Модель в формате Keras (.json и .weights) загружается с помощью класса KerasJS.Model. Основные параметры конструктора:

  • filepath — путь к JSON-файлу модели.
  • backend — можно указать 'cpu' или 'webgl' для выполнения в браузере.

Пример инициализации:

const KerasJS = require('keras-js');

const model = new KerasJS.Model({
  filepaths: {
    model: 'model.json',
    weights: 'model_weights.buf'
  },
  backend: 'cpu'
});

await model.ready();

model.ready() возвращает промис, который разрешается после полной загрузки архитектуры и весов модели.

Структура юнит-тестов

Jest использует глобальные функции describe и test (или it) для структурирования тестов. Основные принципы:

  1. Изоляция тестов — каждый тест должен выполняться независимо, без сохранения состояния между ними.
  2. Фиксированные входные данные — для корректной проверки результатов необходимо заранее определить набор входных тензоров.
  3. Проверка точности предсказаний — результаты сравниваются с эталонными значениями с допустимой погрешностью.

Пример структуры теста для Keras.js:

describe('Тестирование модели классификации', () => {
  let model;

  beforeAll(async () => {
    model = new KerasJS.Model({
      filepaths: { model: 'model.json', weights: 'model_weights.buf' },
      backend: 'cpu'
    });
    await model.ready();
  });

  test('Выходное значение для нулевого входа', async () => {
    const inputData = new Float32Array([0, 0, 0, 0]);
    const output = await model.predict({ input: inputData });
    
    const expected = new Float32Array([0.25, 0.25, 0.25, 0.25]);
    for (let i = 0; i < output.input.length; i++) {
      expect(Math.abs(output.input[i] - expected[i])).toBeLessThan(0.01);
    }
  });
});
  • beforeAll обеспечивает однократную загрузку модели перед запуском всех тестов.
  • Точность проверяется через допустимую разницу toBeLessThan(0.01).

Тестирование различных компонентов модели

Юнит-тестирование можно разделить на несколько уровней:

  1. Проверка структуры модели

    • Соответствие количества слоёв ожидаемому.
    • Проверка форм входов и выходов (inputShape, outputShape).
  2. Тестирование предсказаний

    • Проверка на фиксированных входах с известными результатами.
    • Использование генераторов случайных входов для проверки стабильности предсказаний.
  3. Проверка совместимости бэкендов

    • Запуск модели на CPU и WebGL и сравнение результатов.
    • Проверка консистентности вычислений при использовании Float32Array и Tensor.

Пример проверки формы выходного тензора:

test('Форма выхода модели', async () => {
  const inputData = new Float32Array([1, 2, 3, 4]);
  const output = await model.predict({ input: inputData });
  
  expect(output.input.length).toBe(4);
});

Работа с асинхронными тестами

Так как Keras.js выполняет вычисления асинхронно, все тесты должны быть промис-ориентированными или использовать async/await. Пример:

test('Асинхронная проверка предсказания', async () => {
  const inputData = new Float32Array([0.5, 0.5, 0.5, 0.5]);
  const output = await model.predict({ input: inputData });

  expect(output.input.reduce((a, b) => a + b, 0)).toBeCloseTo(1.0, 5);
});
  • Используется toBeCloseTo для сравнения чисел с плавающей точкой.
  • reduce суммирует элементы выходного массива, что удобно для проверки нормализации вероятностей.

Мокирование и тестирование ошибок

Для проверки устойчивости модели к некорректным данным можно использовать мокированные входные данные или специально подставленные undefined, null, неверные размеры массива. Пример проверки ошибки:

test('Ошибка при неверной форме входа', async () => {
  const invalidInput = new Float32Array([1, 2]); // Ожидается длина 4
  await expect(model.predict({ input: invalidInput }))
    .rejects
    .toThrow(/Input shape mismatch/);
});
  • Используется rejects.toThrow для асинхронного тестирования исключений.

Параметризация тестов

Jest поддерживает запуск одинаковых тестов с разными входными данными с помощью test.each:

test.each([
  [new Float32Array([0, 0, 0, 0]), [0.25, 0.25, 0.25, 0.25]],
  [new Float32Array([1, 1, 1, 1]), [0.25, 0.25, 0.25, 0.25]]
])('Проверка выхода для %p', async (inputData, expected) => {
  const output = await model.predict({ input: inputData });
  for (let i = 0; i < output.input.length; i++) {
    expect(Math.abs(output.input[i] - expected[i])).toBeLessThan(0.01);
  }
});
  • Позволяет компактно описывать множество сценариев.
  • Облегчает поддержку и расширение набора тестов при изменении модели.

Логирование и отладка

Для глубокого анализа работы модели можно включать логирование:

model.predict({ input: inputData })
  .then(output => console.log('Output:', output))
  .catch(err => console.error('Error:', err));
  • Полезно при проверке неожиданных результатов.
  • Позволяет выявить проблемы с весами, формами данных и асинхронной обработкой.

Интеграция с CI/CD

Jest легко интегрируется в пайплайны CI/CD (GitHub Actions, GitLab CI, Jenkins). Основные моменты:

  • Автоматический запуск тестов при каждом коммите.
  • Проверка корректности модели перед деплоем.
  • Возможность генерации отчетов в формате JSON, HTML.
jobs:
  test:
    runs-on: ubuntu-latest
    steps:
      - uses: actions/checkout@v3
      - uses: actions/setup-node@v3
        with:
          node-version: '18'
      - run: npm install
      - run: npm test
  • Обеспечивает контроль качества моделей Keras.js на протяжении всего жизненного цикла проекта.