# Apple MLX: потеря масштабной инвариантности производных arctan2

Локальный технический отчёт, 10 сентября 2026 года.
Дефект воспроизведён на MLX 0.32.2 и настоящем C++ runtime, CPU.
Исследовательский патч прошёл 1368 сравнений; неуспешные промежуточные
варианты сохранены. Это один дефект вычисления производных arctan2,
а не 772 самостоятельные находки.

## Главный контрпример

Для любого C>0 функция

    f(t) = atan2(C·t, C) = atan(t)

не зависит от C. При t=1 её первые три производные в точности равны
0.5, −0.5, 0.5.

| dtype | C | MLX: f′(1) | MLX: f″(1) | MLX: f‴(1) | Центральная разность forward |
|---|---:|---:|---:|---:|---:|
| float32 | 1 | 0.5 | −0.5 | 0.5 | 0.5000209808 |
| float32 | 2⁻⁸⁰ | **inf** | NaN | NaN | 0.5000209808 |
| float32 | 2⁸⁰ | **0** | NaN | NaN | 0.5000209808 |
| float64 | 1 | 0.5 | −0.5 | 0.5 | 0.5000000795 |
| float64 | 2⁻⁶⁰⁰ | **inf** | NaN | NaN | 0.5000000795 |
| float64 | 2⁶⁰⁰ | **0** | NaN | NaN | 0.5000000795 |

Прямой результат MLX во всех строках корректен с точностью dtype:
около π/4. Центральные разности вычислены по значениям **настоящего
MLX forward**, с шагом 2⁻⁶ для float32 и 2⁻¹⁰ для float64.
Они независимо подтверждают первый градиент; второй и третий порядки
проверяются аналитическими производными atan(t).

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

    import mlx.core as mx
    mx.set_default_device(mx.cpu)
    t = mx.array(1., dtype=mx.float32)
    for power in (0, -80, 80):
        c = mx.array(2.**power, dtype=mx.float32)
        f = lambda z: mx.arctan2(c*z, c)
        print(power, f(t).item(), mx.grad(f)(t).item())
    # одинаковый forward, но градиенты: 0.5, inf, 0

[Python-пробник](probe.py), [все результаты wheel](wheel-reproduction.json).
Результат не объясняется особой точкой (0,0): аргументы ненулевые,
а составная функция равна гладкой atan(t) независимо от C.

## Первый порядок и геометрические инварианты

В точке (y,x)=(C,C) градиент atan2(y,x) равен [1/(2C),−1/(2C)].
JVP в радиальном направлении (C,C) должен быть 0, в направлении
(C,−C) — 1.

В float16 при C=1024 исходный VJP возвращает [0,−0]
вместо [0.00048828125,−0.00048828125].
При C=2⁻¹⁴ возвращается [inf,−inf] вместо [8192,−8192].
Оба ожидаемых результата представлены в нормальном диапазоне float16.
Радиальный JVP в обоих случаях равен NaN вместо 0.
Аналогичные нарушения воспроизведены для bfloat16, float32 и float64.

Для обычных эталонов применяется std::hypot(y,x), затем
(x/r)/r и −(y/r)/r; квадрат большого или малого радиуса не формируется.
Для отдельных случаев с очень разными аргументами и весами эталон
вычисляется непосредственно как степень двойки. Пренебрегаемая
поправка квадрата отношения меньше точности эталона и допуска.

## Причина и вариант исправления

Публичные
[ArcTan2::vjp](https://github.com/ml-explore/mlx/blob/81ba1c6a0e50a9268b931579c2d4f1158b9aab5a/mlx/primitives.cpp#L524)
и
[ArcTan2::jvp](https://github.com/ml-explore/mlx/blob/81ba1c6a0e50a9268b931579c2d4f1158b9aab5a/mlx/primitives.cpp#L552)
используют y²+x² непосредственно. Этот знаменатель может стать inf или 0
при конечных аргументах и представимых производных. JVP дополнительно
формирует ненормированные произведения координат и направлений.

[Локальный патч](arctan2-scale-autodiff.patch) выбирает
s=max(|y|,|x|) и D=(y/s)²+(x/s)². При конечной ненулевой паре
1≤D≤2. Для каждого вклада вычисляется g·z/(s²D), где g — входящий
вес, z=x для производной по y и z=−y для производной по x.

Порядок операций выбирается по масштабу и представимости g/s:

| Условие | Вычисление вклада |
|---|---|
| s≥1 | ((g/s)·z/D)/s |
| 0<s<1, g/s конечен | ((g/s)·(z/s))/D |
| 0<s<1, g/s переполняется | ((g·z/D)/s)/s |

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

Масштаб s помечен stop_gradient. Тождество остаётся верным для любого
фиксированного положительного s; значит, его можно дифференцировать
при выбранном постоянном масштабе в окрестности точки.
Выбор порядка по весу также переключает тождественные формулы.
Это сохраняет математическую зависимость от входов и весов.
Второй, третий порядок и смешанный гессиан проверены отдельно.

Это исследовательский вариант корректности с дополнительными операциями.
Его производительность и пригодность для включения в MLX не установлены.

## Почему первой нормировки было недостаточно

Первый патч нормировал сами частные производные, затем умножал на вес.
Он прошёл 1350 проверок, но потерял малую координату в следующем случае:

    float32: y=2⁻¹²⁰, x=2³², g=2¹²⁰
    правильный взвешенный градиент по x: ≈−2⁻⁶⁴
    исходный MLX:                          −2⁻⁶⁴
    первая нормировка:                     0

Отношение y/s=2⁻¹⁵² не представимо в float32; большой вес должен
участвовать до этой потери. Аналогичные случаи найдены для bfloat16
и float64. Это настоящая регрессия первоначального патча.

[Первоначальные исходники и результаты](initial-normalization/).
В расширенной серии 1362 проверок этот вариант даёт 6 несовпадений.
Все шесть соответствующих случаев исходный MLX проходит.

Следующий вариант учитывал вес, но ещё терял результат, если при малом
масштабе умножить два малых числа до завершающего деления:

    float32: y=2⁻⁸⁰, x=2⁻¹²⁰, g=2⁻¹²⁰
    правильный взвешенный градиент по y: ≈2⁻⁸⁰
    промежуточный вариант:                 0

[Второй вариант и его проверки](weighted-two-order/).
Он проходит 1362 проверки, но на расширении до 1368 даёт 6 несовпадений.
Это оставшийся дефект, а не ухудшение относительно исходного MLX,
который также не проходит эти случаи.
Окончательный выбор из трёх порядков устраняет обе группы.

## Итог нативных проверок

| dtype | Сравнения | Несовпадения до | Несовпадения после |
|---|---:|---:|---:|
| float16 | 294 | 162 | 0 |
| bfloat16 | 300 | 166 | 0 |
| float32 | 387 | 222 | 0 |
| float64 | 387 | 222 | 0 |
| Всего | **1368** | **772** | **0** |

- Первый порядок JVP/VJP: четыре dtype, три масштаба, семь направлений
  точки во всех квадрантах и на гладких участках осей; веса 0,1,−0.5.
- Раздельное дифференцирование по каждому аргументу и совместный JVP.
- Радиальная и угловая производные, включая малые и большие масштабы.
- Для float32/float64: производные композиции до третьего порядка
  в семи точках, смешанный гессиан, матрицы 2×3 и неплотные входы.
- Broadcasting (2,1) с (1,3), включая суммирование VJP до исходной формы.
- Дополнительные большие и малые веса при сильно разных аргументах
  для bfloat16, float32 и float64.

Допуски для ненулевого эталона относительные: 5e−13 в float64,
2e−5 в float32, 3e−3 в float16, 2e−2 в bfloat16.
Нулевой эталон проверяется с таким же численным абсолютным допуском.
Обнуление представимой ненулевой производной даёт относительную
ошибку 1 и не может скрыться за этими допусками.

[Нативный тест](native_regression.cpp), [до](run-before.json),
[после](run-after.json), [скрипт](build_and_test.py),
[команды и время](build-results.json).
Серия выполняет настоящие операции MLX, а не их имитацию на Python.
Допуски первоначальной серии не повышались для прохождения исправления.

## Версии, дубликаты и область применимости

- Установленный wheel: MLX 0.32.2, CPU.
- Нативная база: ce916dbbcaa88e433b6fd1e60a17f766d49c27fe.
- Публичный main: 81ba1c6a0e50a9268b931579c2d4f1158b9aab5a.
- Целевые JVP/VJP базы совпадают с сохранённым main.
- [Публичные источники и SHA-256](source-metadata.json),
  [нативные зависимости](native-source-metadata.json).

Поиск выполнен четырьмя запросами GitHub API. Получено 11 уникальных
issues/PR, ответы полные. [Сохранённые результаты поиска](duplicate-search.json).
В частности:

- [#2451](https://github.com/ml-explore/mlx/issues/2451) и
  [#2453](https://github.com/ml-explore/mlx/pull/2453) — прежняя неверная
  формула/число выходных градиентов на умеренных аргументах.
- [#3633](https://github.com/ml-explore/mlx/pull/3633) — индексация
  направлений при частично отслеживаемых входах.
- [#3738](https://github.com/ml-explore/mlx/pull/3738) — повторное
  использование старых записей карты JVP.
- [#4227](https://github.com/ml-explore/mlx/pull/4227) — предложение
  эталонных тестов на умеренных входах.
- [#4257](https://github.com/ml-explore/mlx/pull/4257) — отклонение
  комплексного dtype; arctan2 официально не поддерживает complex64.
- Остальные описания относятся к добавлению операции, алиасам,
  унификации backend и JIT-компиляции Float16.

В прочитанных описаниях точного совпадения с потерей масштабной
инвариантности не найдено. Это ограниченный поиск, не доказательство
приоритета; все комментарии и внешние обсуждения не проверялись.
Security-impact и право на выплату не установлены.

Проверены CPU и указанные вещественные dtype. GPU, compile, vmap,
производительность, все субнормальные режимы, NaN/±inf и особая точка
(0,0) не входят в проверенную область. Нельзя считать конечную серию
доказательством устойчивости при всех комбинациях входов и направлений.

## Нагрузка и сохранность рабочей копии

Работа выполнена последовательно, численные потоки ограничены одним.
GPU-вычислений не было. Импорту установленного wheel потребовалась
системная инициализация Metal; затем явно выбран CPU. Нативный архив
собран без Metal.

Использован существующий CPU-архив и отдельно скомпилированный
primitives.cpp перед ним при линковке. Полной чистой сборки текущего
main не было. Скрипт требует уже имеющихся зависимостей по сохранённым
путям; рабочая копия MLX с прежними изменениями не редактировалась.

Последние два нативных тестовых запуска суммарно заняли около
0.380 секунды CPU. Все пять сохранённых этапов вместе с компиляциями
и линковкой — около 13.43 секунды CPU. Wheel-пробник — около 0.0053
секунды CPU. Эти времена не являются сравнительным benchmark патча.

[Проверка целостности](validate_artifacts.py),
[её результаты](artifact-check-results.json).
Примечание: исходный локальный отчёт до публикации. Актуальное оформление и статус приведены в README.md.
