Zbytečný převod na 32 bitů bral v knihovně TRL přes pětinu času na grafické kartě
Vnitřní smyčka výchozí ztrátové funkce v knihovně TRL převáděla obě vstupní matice ze šestnácti bitů na dvaatřicet, ačkoli obě v paměti ležely v šestnácti. Násobení se tím přesunulo z tensor cores na pomalejší jednotky a v paměti vznikla kopie celé výstupní tabulky o velikosti 2,03 GB. Vydání 1.13.0 ten převod vypustilo.
Knihovna TRL od Hugging Face je sada trenérů k doladění jazykových modelů. Její výchozí ztrátová funkce se od začátku září počítá jinak než dřív a rozdíl je jediný řádek Pythonu. Ven vyšel ve vydání 1.13.0 10. září.

Proč se poslední vrstva počítá po kusech
Poslední vrstva jazykového modelu, lm_head, převádí skrytý stav na skóre pro každý token slovníku. Při dávce 16 384 tokenů a slovníku o 248 320 položkách je výsledkem matice o víc než čtyřech miliardách čísel; ve dvaatřicetibitové přesnosti to je přes 16 GB samotných logitů. Většina karet tolik volné paměti nemá.
Proto se vstup rozseká na kusy a pro každý se projekce, log-softmax i ztráta počítají zvlášť. Celá matice logitů tak nikdy nevznikne naráz. Tahle cesta se v TRL jmenuje chunked_nll a je výchozí.
Převod, který nic nepřinesl
Až do verze 1.12.0 stál ve smyčce tenhle řádek:
logits = h.float() @ w.float().t()
h je skrytý stav, w je váha lm_head. Obojí model drží v bf16, tedy v šestnácti bitech. Volání .float() z nich udělá dvaatřicetibitové kopie.
Podle popisu návrhu změny #6863 tím výpočet nezíská žádnou informaci navíc, protože přesnost vstupů zůstává bf16. Zaplatí ale dvakrát. Násobení matic se přesune z tensor cores, tedy z jednotek, které NVIDIA do svých čipů dala právě kvůli násobení matic, na obyčejná výpočetní jádra. A v paměti vznikne dvaatřicetibitová kopie celé váhy lm_head – znovu pro každý kus a znovu při každém přepočtu v rámci gradient checkpointingu.
Návrh tu kopii vyčísluje na 2,03 GB u slovníku o 248 320 tokenech. Vychází to: 248 320 × 2 048 × 4 bajty je 2 034 237 440 bajtů. Ta dvojice čísel navíc není vzata z ničeho – slovník 248 320 a šířku 2 048 má podle souboru config.json model Qwen3.6-35B-A3B, na kterém autor pouštěl profil.
985 milisekund ze 4,57sekundového kroku
Všechna čísla níž pocházejí z toho návrhu a nikdo nezávislý je nezopakoval. Měřil je Quentin Gallouédec z Hugging Face na osmi kartách H100: trl sft nad Qwen3.6-35B-A3B, FSDP2, délka sekvence 4 096, dávka čtyři na kartu, LoRA. Dvě jádra pro dvaatřicetibitové násobení matic v tom profilu zabrala 985 ms ze 4,57sekundového kroku, tedy 21,6 % veškerého času stráveného na GPU, ve 192 spuštěních po průměrných 5,1 ms. Obě čísla spolu drží: 192 × 5,1 ms je 979 ms a 985 ze 4 570 dá 21,55 %.
Samotný kus – 256 tokenů, slovník 248 320, šířka 2 048, jedna H100, dopředný i zpětný průchod – trval 23,37 ms a zabral 5,99 GB. Po změně 3,86 ms a 3,03 GB. To je šestinásobek rychlosti a o 2,96 GB méně paměti.
Oprava je jeden řádek
Návrh přidal 28 řádků a pět smazal ve čtyřech souborech, z toho dva jsou testy. V trl/trainer/sft_trainer.py se změnil tenhle jediný:
logits = (h @ w.to(h.dtype).t()).float()
Násobení proběhne v datovém typu modelu a na dvaatřicet bitů se převede až výsledek, kvůli softmaxu. Přesně to dělá i alternativní loss_type="nll" a funkce ForCausalLMLoss v knihovně Transformers. Že oprava opravdu je ve vydání, jde ověřit přímo: ve značce v1.12.0 stojí na 103. řádku starý tvar, ve v1.13.0 na 105. řádku nový.
Numerika se podle autora pod smíšenou přesností knihovny accelerate nemění ani o bit, protože autocast oba operandy stejně převáděl zpátky dolů. Destilační trenér platil týž převod dvakrát na kus, jednou za studenta a jednou za učitele; opravený je taky, změřený ale ne.
Na pěti modelech od 1,20násobku do 1,69násobku
Průchodnost celého trénování v tokenech za sekundu na kartu, 16 384 tokenů na krok, všude s chunked_nll:
| model | režim | před | po | |
|---|---|---|---|---|
| gemma-3-270m (slovník 262k) | plné doladění, 1× H100 | 24 036 | 31 609 | 1,32× |
| Qwen3-0.6B (slovník 152k) | plné doladění, 1× H100 | 21 276 | 26 641 | 1,25× |
| Qwen3-8B | plné doladění, 2× H100 FSDP2 | 3 554 | 6 009 | 1,69× |
| Qwen3-8B | LoRA r16, 2× H100 FSDP2 | 4 531 | 7 125 | 1,57× |
| Qwen3-30B-A3B | LoRA na pozornosti, 2× H100 FSDP2 | 4 251 | 5 101 | 1,20× |
Šestinásobek na jednom kusu a 1,69násobek na celém trénování si neodporují. Zrychlené jádro musí prorazit vším ostatním, co v kroku běží, a čím menší podíl na něm poslední vrstva má, tím míň je změna vidět. Nejhůř dopadl model se směsí expertů, kde se LoRA dotýká jen pozornosti. Že se zrychlené jádro nemusí projevit vůbec, ukázal SGLang u modelu Qwen-Image.
Kde převod zůstal
Dvě místa nechal návrh vědomě být. V souboru trl/trainer/utils.py na 1446. řádku stojí grad_hidden.add_(grad_logits @ w_chunk.float()) – tam je grad_logits opravdu dvaatřicetibitový, takže převod váhy smysl má. A experimentální asynchronní destilační trenér nese třetí kopii; jeho projekce běží mimo accelerator.autocast(), takže by šestnáctibitové násobení chtělo převod váhy napsat výslovně, a to autor v experimentálním kódu dělat nechtěl.
Verze 1.13.0 je na PyPI od 10. září, předchozí 1.12.0 tam leží od 26. srpna. Kdo trénuje s výchozím nastavením a na starší verzi, platí ten převod dál.
Zdroje
- Návrh změny huggingface/trl#6863 včetně měření a rozboru
- Poznámky k vydání TRL 1.13.0
- Balíček trl 1.13.0 na PyPI
- Nastavení modelu Qwen3.6-35B-A3B
- GPU Performance Background User's Guide od NVIDIA