Distributed group_mean casts group ids to float32, losing precision vs non-distributed path
メンテナーはふだん 1 日以内に返信
評価
- 難易度
- 2/5
- 見積もり時間
- 1〜3時間
- 初心者へのやさしさ
- 82/100
- issue の種類
- バグ
- 明瞭さ
- 明確に書かれている
- 活発さ
- 活発
調査の方向性
src/art/megatron/context_parallel/loss_inputs.py の ContextParallelLossInputs._distributed_group_mean から始め、art.utils.group_aggregate.group_aggregate および AlignedLossInputs.group_mean のパスと比較します。unique と searchsorted を通して by のネイティブな dtype を保持し、分散グルーピングが非分散パスと同じ精度保証を維持していることを確認します。
索引モデルが issue の本文から書いたものです。
説明
Summary
`ContextParallelLossInputs._distributed_group_mean` in `src/art/megatron/context_parallel/loss_inputs.py` casts the grouping key (`by`) to `float32` before `torch.unique`/`torch.searchsorted`:
```python
flat_by = by.reshape(-1).to(dtype=torch.float32)
```
The non-distributed path (`art.utils.group_aggregate.group_aggregate`, used by the base `AlignedLossInputs.group_mean`) preserves the tensor's native dtype instead, and explicitly documents that "any dtype accepted by `torch.unique` is supported."
Failure scenario
`float32` cannot represent all integers exactly once they exceed 2**24 (16,777,216). If a grouping key ever carries a value past that range (e.g. a UID-based or otherwise large-range grouping key), two distinct group ids can collapse to the same float32 value, or get mis-bucketed by `searchsorted` against the globally-gathered sorted ids. This would silently merge two groups' values into the same mean/denominator on context-parallel ranks, corrupting the loss without raising any error.
This is low-probability today since group ids used in practice are small, but it's a real, silent discrepancy between the distributed and non-distributed code paths that should use the same precision guarantees.
Fix
Keep `by`'s native dtype through `unique`/`searchsorted` instead of forcing a float32 round-trip (matching `group_aggregate`'s contract). PR incoming.
- 主要言語
- Python
- スター
- 10.8k
- フォーク
- 1k
- 平均マージ
- 15時間 31分
- マージ済み PR(30日)
- 166
環境構築
- Dockerfile・Docker Compose ファイルなし
- プルリクエストのテンプレートなし
- コントリビューションガイドを読む
はじめの一歩
- issue を最後まで読み、次にプロジェクトのコントリビューションガイドを読みます。
- 着手することを issue にコメントします — 二人が同じ作業をするのを防げます。
- リポジトリをフォークし、ブランチを切って変更します。
- issue 番号を参照したプルリクエストを送ります。
OpenPipe/ART のほかの issue
-
Add from_entity parameter to _experimental_fork_checkpoint再び着手できるかも このイシューのプルリクエストはマージされずにクローズされました。 オープン
難易度 2/5 1〜3時間 初心者へのやさしさ 72/100
メンテナーはふだん 1 日以内に返信
-
Nonfused column LoRA misses input-gradient SUM when TP>1 and sequence parallelism is disabled対応中かも このイシューにリンクされたプルリクエストがオープン中、またはマージ済みです。 オープン
難易度 5/5 1週間以上 初心者へのやさしさ 25/100
メンテナーはふだん 1 日以内に返信
-
難易度 4/5 3〜5日 初心者へのやさしさ 54/100
メンテナーはふだん 1 日以内に返信
-
難易度 5/5 1週間以上 初心者へのやさしさ 25/100
メンテナーはふだん 1 日以内に返信
-
難易度 5/5 1週間以上 初心者へのやさしさ 10/100
メンテナーはふだん 1 日以内に返信
似ている issue
-
Harmony OPeNDAP SubSetter (HOSS) Geographic LARC_CLOUD PREFIRE_SAT2_AUX-SAT R01 production
難易度 2/5 1〜3時間 初心者へのやさしさ 68/100
nasa/harmony-autotester#245 ·
-
feature
難易度 2/5 1〜3時間 初心者へのやさしさ 66/100
-
L: github:actions L: php:composer
難易度 2/5 1〜3時間 初心者へのやさしさ 88/100
dependabot/dependabot-core#16493 ·
メンテナーはふだん 1 日以内に返信
-
難易度 1/5 1時間未満 初心者へのやさしさ 92/100
DataTalksClub/machine-learning-zoomcamp#730 ·
メンテナーはふだん 2 日以内に返信
-
難易度 2/5 1〜3時間 初心者へのやさしさ 82/100
メンテナーはふだん 4 日以内に返信