# Apple MLX: FP16 BatchNorm повреждает накопленную дисперсию

Локальный технический отчёт, 10 сентября 2026 года. Ошибка воспроизведена
на настоящем MLX 0.32.2, CPU. Подготовлен изолированный патч расчёта
статистики BatchNorm. Ничего не опубликовано, не отправлено и не применено
к рабочей копии MLX.

## Подтверждённое поведение

Одного конечного FP16-пакета достаточно, чтобы `running_var`, хранящаяся
по умолчанию в float32, стала бесконечной. Последующие небольшие пакеты
не восстанавливают это состояние. В режиме eval слой выдаёт нули
вместо конечного ненулевого результата.

Минимальное воспроизведение:

```python
import mlx.core as mx
import mlx.nn as nn
mx.set_default_device(mx.cpu)
bn = nn.BatchNorm(1, momentum=0.125, affine=False)

x = mx.array([[-256.], [256.]], dtype=mx.float16)
print(bn(x).tolist())             # [[-0.0], [0.0]] вместо примерно [-1,1]
print(bn.running_var.tolist())   # [inf] вместо [16384.875]

small = mx.array([[-1.], [1.]], dtype=mx.float16)
bn(small)
print(bn.running_var.tolist())   # всё ещё [inf], правильно [14337.015625]
bn.eval()
print(bn(small).tolist())        # [[-0.0], [0.0]] вместо примерно ±0.008351618
```

| Этап | Исходный MLX | После патча |
|---|---|---|
| train на `[-256,256]` | `[0,0]` | `[-1,1]` |
| running_var после первого пакета | `inf` | 16384.875 |
| running_var после малого `[-1,1]` | `inf` | 14337.015625 |
| eval на `[-1,1]` | `[0,0]` | ≈`[-0.008351618,0.008351618]` |

Числа таблицы получены настоящими вызовами исходного/исправленного слоя:
[compatibility-results.json](compatibility-results.json), записи `momentum/0.125/*`.
[Первый пробник](probe.py) отдельно проверяет три последующих малых пакета
и четыре dtype; [результат установленного wheel](wheel-reproduction.json).
Контрпример не требует уже существовавших NaN/Inf или непредставимых входов.

## Причина

Публичный исходник на момент чтения:
`81ba1c6a0e50a9268b931579c2d4f1158b9aab5a`.
В [BatchNorm._calc_stats](https://github.com/ml-explore/mlx/blob/81ba1c6a0e50a9268b931579c2d4f1158b9aab5a/python/mlx/nn/layers/normalization.py#L355)
mean и var вычисляются непосредственно в dtype входа. Приведение к типу
накопленного состояния происходит уже после потери результата.

Для двух чисел −C,C среднее равно 0, дисперсия для train равна C²,
а несмещённая оценка для running_var равна 2C². При C=256 это 65536 и
131072: промежуточные значения не помещаются в FP16, хотя обновлённая
float32-статистика должна быть конечной:

    v₁ = (1−1/8)·1 + (1/8)·131072 = 16384.875.

При следующих конечных пакетах множитель 7/8 сохраняет уже возникшую
бесконечность. Обучение на небольшом пакете снова может выдавать нормальный
результат, поскольку использует его собственную дисперсию. Это не означает,
что восстановилась накопленная статистика для eval.

Отдельный граничный случай `[-192,192]` показывает более раннее переполнение:
настоящая train-дисперсия **36864 представима в FP16**, но сумма двух
квадратов 36864+36864 переполняется до деления на 2.
Это соответствует порядку операций в
[mx.var](https://github.com/ml-explore/mlx/blob/81ba1c6a0e50a9268b931579c2d4f1158b9aab5a/mlx/ops.cpp#L2453).
Записаны реальные отдельные квадраты, их сумма, дисперсия и состояние
[до/после](variance-boundary-results.json); [скрипт](variance_boundary.py).

Проверены и обучаемые параметры: для трёхэлементного FP16-пакета с масштабом
256 исходный `nn.value_and_grad` возвращает градиент gamma `[0,0]` вместо
примерно `[-1.3363062,-0.7349684]`. Исправленный слой проходит независимый
аналитический эталон для gamma, beta и входа. Накопленные статистики
остаются замороженными параметрами, но обновляются в value_and_grad.

## Исправление

[Патч](batchnorm-fp32-statistics.patch) добавляет семь строк в два метода
BatchNorm:

- FP16/BF16 вход приводится к float32 перед расчётом mean и var.
- Нормализация использует эти статистики. В train и при
  `track_running_stats=False` её результат возвращается в исходный
  низкоточный dtype до применения affine-параметров.

Отдельные оценки дисперсии для train (`ddof=0`) и running_var (`ddof=1`),
правило momentum, проверка размера пакета и хранение замороженных параметров
сохранены. Float32/float64 расчёт не изменён. В eval с накопленной статистикой
сохранено прежнее продвижение dtype, включая float32-выход при FP16-входе
и стандартных float32-буферах.

Если пользователь предварительно привёл буферы running_mean/running_var
к FP16/BF16, следующий train-вызов прототипа повышает их до float32.
Это сознательная политика для статистики, отдельно проверенная в тесте.
Она требует согласования при ревью. Патч предотвращает новое повреждение;
**уже сохранённые бесконечные статистики он не восстанавливает**.

Тесты загружают точные сохранённые Python-модули до и после патча и выполняют
их на установленном бинарном MLX 0.32.2. Это не полная сборка текущего main.
Текст класса установленного wheel также сохранён и сверен с публичным.
[Версия, исходники, SHA256](source-metadata.json),
[установленный класс](installed-batchnorm-source.py),
[построение патча](build_patch.py), [исправленный модуль](patched-normalization.py).

## Результаты проверки

[Основной набор](regression.py): 192 сценария forward/state/eval и 12
сценариев с градиентами. Эталон средних и дисперсий рассчитывается скалярной
Decimal-арифметикой с точностью 60 цифр по уже квантованным входам.
Производные сравниваются с аналитической формулой BatchNorm. Округление
эталона учитывает dtype нормализованного выхода и dtype параметров.

| Набор | Проверок до / после | Ошибок до | Ошибок после |
|---|---:|---:|---:|
| Основной | 2760 / 2760 | 102 | 0 |
| Совместимость и крайние параметры | 29 / 33 | 18 | 0 |
| Граница с представимой train-дисперсией | 5 / 5 | 4 | 0 |

2760 — число assertions, не независимых пакетов. Из них 1020 сравнивают
численные результаты, остальные — типы, форму, неизменность входов,
состояния eval и отсутствие градиентов замороженной статистики.
Основные 102 несовпадения распределены так: 68 выходов, 24 running_var,
4 running_mean, 2 gamma-градиента, 2 входных градиента, 2 состояния
после value_and_grad. 86 относятся к float16, 16 — к bfloat16.
В счёт входят также ошибки точности низкоточной статистики; не все
несовпадения являются переполнениями.

Проверены NC, NLC, NHWC; четыре dtype; малые, обычные и большие масштабы;
affine on/off; tracking on/off; train→eval; неодинаковые gamma/beta.

[Совместимость](compatibility.py) включает точные публичные методы
`test_batch_norm` и `test_batch_norm_stats` с настоящими unittest-assertions,
momentum 0/0.1/0.125/1, положительные eps вплоть до 10⁻¹² на постоянном
входе, ошибочные формы, сохранение и загрузку маленьких локальных NPZ
checkpoint. Четыре дополнительные проверки после относятся к сознательному
повышению типа предварительно приведённых буферов.
PyTorch-parity метод не запускался; независимый эталон описан выше.

Результаты: [основной набор](regression-results.json),
[совместимость](compatibility-results.json).
CPU-время последних численных запусков: приблизительно 0.667, 0.024 и
0.008 секунды. Это длительность тестов, не benchmark самого слоя.
Численные потоки установлены в 1, все операции явно на CPU и последовательно.
GPU, распределённое обучение, большие модели и полная сборка не запускались.

## Связь с известными работами

[Поиск](duplicate-search.json) вернул 26 уникальных issues/PR. Содержательно
проверены ближайшие обсуждения BatchNorm и нормализации, включая #3817,
#1960, #217, #4228; сохранён и прочитан 31 комментарий из трёх issue/PR
веток и review #3817: [комментарии](duplicate-comments.json).

- [#3817](https://github.com/ml-explore/mlx/pull/3817) исправляет несмещённую
  running_var и одиночный training-пакет. Это исправление уже присутствует;
  новый патч не меняет ddof и не повторяет ту находку.
- [#1960](https://github.com/ml-explore/mlx/issues/1960) обсуждает момент
  обновления состояния в value_and_grad. Здесь состояние обновляется,
  но получает бесконечное значение.
- [#4228](https://github.com/ml-explore/mlx/issues/4228) описывает известное
  FP16-переполнение InstanceNorm. Общий численный механизм родственен.
  Нынешний пакет воспроизводит его в BatchNorm и показывает сохранение
  повреждённой статистики между train-вызовами и в eval.
- [#217](https://github.com/ml-explore/mlx/pull/217) — исходное добавление
  BatchNorm и обсуждение интерфейса/оценки дисперсии.

Точного отчёта о данном FP16 BatchNorm-контрпримере в просмотренных
материалах не найдено. Это ограниченный поиск, без гарантии абсолютной
новизны и без утверждения о новом классе уязвимостей безопасности.

## Оставшиеся ограничения

При входах порядка 2⁸⁰ у float32/BF16 сама дисперсия превышает float32.
Одна смена dtype с BF16 на float32 этого не устраняет: дополнительная
проверка после патча снова даёт нули и inf в таких режимах. Они явно
сохранены в `remaining_limits` файла compatibility-results.json и не
выданы за исправленные случаи.

GroupNorm в первом пробнике также возвращает нули для конечного FP16-входа
в стандартном режиме, тогда как `pytorch_compatible=True` при одной группе
возвращает ±1. Это соседнее измерение; патч GroupNorm, полный аудит его
дубликатов и отдельный завершённый пакет здесь не подготовлены.
Самостоятельные `mx.mean`/`mx.var` данным патчем тоже не изменяются.

Compiled-путь, GPU, полное обучение модели и производительность на больших
массивах не проверены. Восстановление повреждённого checkpoint и финансовый
ущерб не продемонстрированы. Патч пока исследовательский, для ревью.

## История проверки

Первый основной запуск дал три ложных несовпадения: float64-вход сравнивался
со слишком точным эталоном градиента gamma, хотя gamma-параметр был float32.
Эталон исправлен с учётом типа параметра; код MLX для этого не менялся.
[Первый результат](regression-first-results.json), [заметки](harness-notes.txt).

Предположение, что `[-192,192]` изолирует только переполнение несмещённой
оценки, не подтвердилось: переполнение суммы ломает и training-нормализацию.
Измерения и порядок операций указаны выше; первоначальный скрипт
[сохранён](variance_boundary-first.py).

Первое чтение источника завершилось DNS timeout. Независимый поиск issues
продолжился; повторное чтение источника успешно. Сохранены
[первый ответ](source-metadata-first.json) и окончательные метаданные.
Снимок ops.cpp повторно использован из предыдущего аудита того же точного
публичного commit с проверкой SHA256.

## Воспроизведение локально

Последовательно из этой папки:

```sh
python3 build_patch.py
/Users/khamit/Documents/math/software-audits/current-stack-2026-09-07/.venv/bin/python probe.py
/Users/khamit/Documents/math/software-audits/current-stack-2026-09-07/.venv/bin/python regression.py
/Users/khamit/Documents/math/software-audits/current-stack-2026-09-07/.venv/bin/python compatibility.py
/Users/khamit/Documents/math/software-audits/current-stack-2026-09-07/.venv/bin/python variance_boundary.py
python3 validate_artifacts.py
```

Численные скрипты сами задают CPU-устройство, один поток и лимит 30 CPU-секунд.
Для повторения тестов сеть не требуется. Доступ за пределами песочницы
потребовался для инициализации установленного MLX, до выбора CPU-вычислений.
[Проверка артефактов](artifact-validation.json), [контрольные суммы](SHA256SUMS.json).
