Distributed group_mean casts group ids to float32, losing precision vs non-distributed path
Los mantenedores suelen responder en 1 día
Evaluación
- Dificultad
- 2/5
- Tiempo estimado
- 1-3 horas
- Aptitud para principiantes
- 82/100
- Tipo de issue
- Error
- Claridad
- Bien especificado
- Estado de actividad
- Activo
Línea de trabajo
Comienza en src/art/megatron/context_parallel/loss_inputs.py, en ContextParallelLossInputs._distributed_group_mean, y compáralo después con art.utils.group_aggregate.group_aggregate y con la ruta group_mean de AlignedLossInputs. Conserva el tipo de datos nativo de by a través de unique y searchsorted, y verifica que la agrupación distribuida mantenga las mismas garantías de precisión que la ruta no distribuida.
Escrito por el modelo de indexación a partir del texto del issue.
Descripción
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.
- Lenguaje dominante
- Python
- Estrellas
- 10.8k
- Forks
- 1k
- Merge medio
- 15 h 31 min
- PR fusionados (30 d)
- 166
Preparar el entorno
- Sin Dockerfile ni archivo de Docker Compose
- Sin plantilla de pull request
- Leer la guía de contribución
Primeros pasos
- Lee el issue completo y luego la guía de contribución del proyecto.
- Comenta en el issue que vas a ocuparte — evita que dos personas hagan lo mismo.
- Haz un fork del repositorio y trabaja en una rama.
- Abre un pull request que haga referencia al número del issue.
Más de OpenPipe/ART
-
Add from_entity parameter to _experimental_fork_checkpointQuizá libre de nuevo Un pull request para esta issue se cerró sin fusionarse. Abierto
Dificultad 2/5 1-3 horas Aptitud para principiantes 72/100
Los mantenedores suelen responder en 1 día
-
Nonfused column LoRA misses input-gradient SUM when TP>1 and sequence parallelism is disabledPosiblemente ocupada Un pull request vinculado a esta issue está abierto o ya se fusionó. Abierto
Dificultad 5/5 Más de una semana Aptitud para principiantes 25/100
Los mantenedores suelen responder en 1 día
-
Dificultad 4/5 3-5 días Aptitud para principiantes 54/100
OpenPipe/ART#961 · 3 comentarios ·
Los mantenedores suelen responder en 1 día
-
Dificultad 5/5 Más de una semana Aptitud para principiantes 25/100
OpenPipe/ART#949 · 5 comentarios ·
Los mantenedores suelen responder en 1 día
-
Dificultad 5/5 Más de una semana Aptitud para principiantes 10/100
Los mantenedores suelen responder en 1 día
Todos los issues de OpenPipe/ART
Issues similares
-
Dificultad 2/5 1-3 horas Aptitud para principiantes 83/100
Los mantenedores suelen responder en 1 día
-
Dificultad 2/5 1-3 horas Aptitud para principiantes 86/100
FuRongJun-1999/dsh-memory#65 ·
Los mantenedores suelen responder en 1 día
-
ci needs-ac
Dificultad 2/5 1-3 horas Aptitud para principiantes 75/100
Ikalus1988/MisakaNet#2930 ·
Los mantenedores suelen responder en 1 día
-
`FakeBackendV2.run` fails with `NoiseError` on circuits with delays on qubits where T2 > 2·T1Abiertobug
Dificultad 2/5 1-3 horas Aptitud para principiantes 78/100
Qiskit/qiskit-aer#2466 ·
-
area/cli
Dificultad 2/5 1-3 horas Aptitud para principiantes 82/100