選ぶ位置は制御値、選ぶ値は計算対象
MLX-LM v0.32.0 に含まれる PR #1930 は、スコアを取り出す前に、argpartition が返す位置へ mx.stop_gradient を適用する修正である。作者は限定した逆伝播テストも追加した。リリースの公開時刻は不明で、PyPIへの配布は別途10月1日と確認できる。
argpartitionの仕様では、戻り値は整数の位置で、区分内の並び順は未定義である。stop_gradientの仕様では、値を変えず、その配列を通る勾配を止める。現在のMLX 0.32.3の説明であり、すべてのインストール環境との互換性を検証した結果ではない。
以下の独自の例では、スコア自体は微分の対象になる。一方、整数の位置は「どの値を取り出すか」を制御する。位置を勾配計算から切り離しても、元の配列から取り出したスコアには、選択が固定された範囲で勾配を流せる。上位k個の境界で同点や順位逆転が起きると、その範囲は変わる。逆伝播の結果が有限で、固定した選択の解析値に合うことを確認し、同点時の勝者や選択が変わる地点の滑らかな微分は求めない。
- 1学習対象のスコア
- 2argpartitionで上位k個を選ぶ
- 3整数の位置
- 4stop_gradient
- 1元のスコア + 切り離した位置
- 2take_along_axis
- 3選ばれたスコア
- 4損失
- 5元のスコアへの勾配
- 1同点または順位逆転
- 2選ばれる位置の集合が変わる
- 3同点時の勝者を固定しない
小さな配列で確かめる未実行の演習
この独自の演習は未実行である。PyPIの情報では、mlx-lm==0.32.0 はPython 3.11以降、DarwinではMLX 0.32.2以降を要求する。Apple siliconでこの経路を使い、実際に解決されたMLXの版を記録する。最低版の指定だけでは依存関係を再現できない。読み込みや版の確認に失敗したら、そこで原因を調べる。
from importlib.metadata import version
import mlx.core as mx
k = 2
def loss(gates):
inds = mx.argpartition(-gates, kth=k - 1, axis=-1)[..., :k]
inds = mx.stop_gradient(inds)
chosen = mx.take_along_axis(gates, inds, axis=-1)
return mx.sum(chosen * chosen)
assert version("mlx-lm") == "0.32.0"
print({"mlx_lm": version("mlx-lm"), "mlx": version("mlx")})
# 形状・型・有限値だけでは、誤って全て0になった勾配も通る。
baseline = mx.array([[1.01, 1.0, 0.2, -0.1]], dtype=mx.float32)
baseline_grad = mx.grad(loss)(baseline)
expected_baseline = mx.array([[2.02, 2.0, 0.0, 0.0]], dtype=mx.float32)
mx.eval(baseline_grad)
assert baseline_grad.shape == baseline.shape
assert baseline_grad.dtype == baseline.dtype
assert bool(mx.all(mx.isfinite(baseline_grad)))
assert bool(mx.allclose(baseline_grad, expected_baseline))
# 境界で同値になる2つのうち、選ばれる位置は固定しない。
boundary = mx.array([[1.0, 0.2, 0.2, -0.1]], dtype=mx.float32)
boundary_grad = mx.grad(loss)(boundary)
mx.eval(boundary_grad)
assert bool(mx.all(mx.isfinite(boundary_grad)))
assert bool(mx.allclose(boundary_grad[:, [0, 3]],
mx.array([[2.0, 0.0]], dtype=mx.float32)))
assert bool(mx.allclose(mx.sum(boundary_grad), mx.array(2.4, dtype=mx.float32)))
assert bool(mx.allclose(mx.sort(boundary_grad[0, 1:3]),
mx.array([0.0, 0.4], dtype=mx.float32)))どちらの配列も形状は (1, 4)、型は float32 である。同点のない基準入力は最初の2つを選ぶ。選ばれたスコアの二乗を微分すると、期待値は [[2.02, 2.0, 0.0, 0.0]] になる。これは例から計算した解析値であり、測定値ではない。形状・型・有限値だけの確認では、誤って全て0になった勾配も通ってしまう。
境界の入力は [1.0, 0.2, 0.2, -0.1]、k=2 である。先頭の勾配は 2.0、末尾は 0.0、合計は 2.4 となる。同点の中間2つは、片方だけが 0.4 になる。その2値を並べ替えて [0.0, 0.4] と照合すれば、勝者や argpartition の順序を固定せずに確認できる。MLXとMLX-LMの版、macOSの版、機器、配列の値、k、例外を記録する。パッケージの不在、版の不一致、機器の非対応、有限でない勾配、仕様の変化は検証失敗として扱う。固定した版の実装を調べ、別の実行基盤へのフォールバックは加えない。
上流のテストは作者が保守する根拠であり、独立した正しさや収束の検証結果ではない。この演習はMistral4やMiniMax全体、最適化器の状態、分散学習、量子化、モデル品質を試験しない。依存関係をインストールした後の計算とメモリは必要だが、有料APIは呼ばない。取得するコードと重みは、信頼性とライセンスを別々に確認する。MLX-LMのMITライセンスが対象にするのはリポジトリのコードであり、全モデルやデータセットの利用条件までは定めない。
2024年の背景と、今回の修正を区別する
Appleの2024年7月17日のICML案内には、Apple silicon上でMLXを使う機器内の推論・学習デモが記されている。この背景はMLX-LM 0.32の性能測定や、今回の修正の根拠にはならない。修正の根拠は2026年のパッチであり、小さな演習は別に用意した未実行の教材である。
MENTAL MODEL / 考える順序
発表から、自分の判断へ。
発表の主張と、論文・公式ドキュメントの条件を並べて読む。
出典
公開日は資料の日付、確認日は内容を参照した日です。コミュニティの観測は公式の確定事項と区別します。