Hacktoberfest 2026:メンテナが10月に向けて印を付けた、オープンで初心者向けの issue。 Hacktoberfest の issue を見る

Distributed group_mean casts group ids to float32, losing precision vs non-distributed path

オープン 初心者向け
#1,102 コメント 0 件 リアクション 0 件 担当者 0 名 GitHub で見る

メンテナーはふだん 1 日以内に返信

@OnePunchMonk がすでに取り組んでいます。

2026年10月3日 から。

  • #1103 @OnePunchMonk による — オープン

評価

難易度
2/5
見積もり時間
1〜3時間
初心者へのやさしさ
82/100
issue の種類
バグ
明瞭さ
明確に書かれている
活発さ
活発
技術スタック
python, pytorch

調査の方向性

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

環境構築

はじめの一歩

  1. issue を最後まで読み、次にプロジェクトのコントリビューションガイドを読みます。
  2. 着手することを issue にコメントします — 二人が同じ作業をするのを防げます。
  3. リポジトリをフォークし、ブランチを切って変更します。
  4. issue 番号を参照したプルリクエストを送ります。

OpenPipe/ART のほかの issue

OpenPipe/ART の issue をすべて見る

似ている issue

Python の issue をもっと見る

新しい issue をメールで受け取る

初心者向けの GitHub issue を短くまとめたダイジェスト。