Přeskočit na obsah
Čísla 3 min čtení

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áří.

Sál superpočítače se stojany výpočetních uzlů
Sál superpočítače Frontier v Národní laboratoři v Oak Ridge. Foto: oakridgelabnews, Flickr (CC BY 2.0)

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:

modelrežimpředpo
gemma-3-270m (slovník 262k)plné doladění, 1× H10024 03631 6091,32×
Qwen3-0.6B (slovník 152k)plné doladění, 1× H10021 27626 6411,25×
Qwen3-8Bplné doladění, 2× H100 FSDP23 5546 0091,69×
Qwen3-8BLoRA r16, 2× H100 FSDP24 5317 1251,57×
Qwen3-30B-A3BLoRA na pozornosti, 2× H100 FSDP24 2515 1011,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

Diskuse

Zatím tu nikdo nediskutuje. Můžete být první.

Napsat příspěvek

Diskutovat můžete i bez účtu. S registrací se ale příspěvek zveřejní hned a nemusíte pokaždé vyplňovat jméno. Účet už máte? Přihlaste se.

Nezveřejňujeme ho, slouží jen redakci.

Podporuje zápis Texy: **tučně**, *kurzíva*, odrážky, odkazy.

Dál k tématu

  1. Čísla

    Globální KV cache DeepSeeku V4.1 Flash zabere při milionu tokenů 890 MiB

    DeepSeek zveřejnil váhy multimodálního modelu V4.1 Flash s kontextem dlouhým 1 048 576 tokenů. Globální KV cache ukládá 890 bajtů na token, takže při zaplnění celého…

  2. Čísla

    Modely Flare a Sunburst z GPT Image 2.5 mají stejné sazby za token

    OpenAI rozdělila GPT Image 2.5 do dvou modelů pro API. Flare má rychleji obsloužit běžné generování, Sunburst má přesněji upravovat obraz; sazby za textové i obrazové…

  3. Čísla

    Pydantic AI 2.40 po zapnutí každou hodinu stahuje 1 678 cenových položek

    Pydantic AI 2.40 přidala volitelné stahování cen modelů hned po startu a potom každou hodinu. Při našem testu se seznam rozšířil z 1 646 na 1 678 položek; výpadek…

  4. Čísla

    Známku ověřeného měření nemá ani jeden z 2 542 výsledků v žebříčcích Hugging Face

    Hugging Face ukládá od prosince 2025 výsledky testů přímo do repozitáře modelu a skládá z nich žebříčky benchmarků. Ze 2 542 položek, které v nich 4. září 2026 stály,…