Kontrollarmen i finjusteringen kjører nå. Det er den vanlige finjusteringen, uten halevekting, som resten av forsøket skal måles mot. For en uke siden var det ikke klart at den kunne kjøre her i det hele tatt.
MET trente Bris med åtte skjermkort per ensemblemedlem. På eX3 er noden med åtte A100-kort tatt ut av drift, og alle kortene i DGX-noden er holdt av jobber som ber om ett kort i en uke av gangen. Det klyngen faktisk gir, pålitelig og i løpet av sekunder, er ett kort. Så spørsmålet ble om modellen kunne presses inn på det.
Hva den kjører på
Dette er jobben mens den trener:
+-----------------------------------------------------------------------------------------+
| NVIDIA-SMI 615.71.09 KMD Version: 615.71.09 CUDA UMD Version: 13.4 |
+-----------------------------------------+------------------------+----------------------+
| GPU Name Persistence-M | Bus-Id Disp.A | Volatile Uncorr. ECC |
| Fan Temp Perf Pwr:Usage/Cap | Memory-Usage | GPU-Util Compute M. |
| | | MIG M. |
|=========================================+========================+======================|
| 0 NVIDIA H200 NVL On | 00000000:4C:00.0 Off | 0 |
| N/A 64C P0 372W / 600W | 143027MiB / 143771MiB | 100% Default |
| | | Disabled |
+-----------------------------------------+------------------------+----------------------+
+-----------------------------------------------------------------------------------------+
| Processes: |
| GPU GI CI PID Type Process name GPU Memory |
| ID ID Usage |
|=========================================================================================|
| 0 N/A N/A 848014 C ...tra/bris-env/.venv/bin/python 14295... |
+-----------------------------------------------------------------------------------------+
Ett H200-kort med 143 771 MiB minne, der treningen bruker 143 027. Det er under én gigabyte igjen. Kortet regner for fullt, og hvert treningssteg tar rundt 33 sekunder, så 3000 steg blir omtrent et døgn.
Hvorfor det ikke passet
Bris jobber på et strukket rutenett. Norden dekkes med 2,5 km fra MEPS, resten av kloden grovere, til sammen 1 359 281 punkter. En koder presser dem ned på et grovere nett med 261 634 noder, en prosessor med 16 transformerlag jobber der, og en dekoder skriver resultatet tilbake til hele rutenettet.
Hvert treningssteg lager flere varsler av samme tidspunkt, med ulik støy, og måler dem mot fasiten som et ensemble. anemoi, rammeverket Bris er bygget på, stabler alle medlemmene i én batch og kjører dem gjennom modellen samtidig. To medlemmer betyr altså at dekoderen jobber på 2,7 millioner punkter på én gang.
Det første forsøket med fire medlemmer manglet 24 GB. To medlemmer manglet 31.
Det som ikke virket
Flere biter i konfigurasjonen. Oppmerksomheten i modellen kan regnes i biter for å spare minne. Jeg økte antallet fra fire til seksten, og allokeringen som feilet var nøyaktig like stor som før, til byten. Grunnen står i en linje i anemoi:
num_chunks = self.num_chunks if self.training else NUM_CHUNKS_INFERENCE_MAPPER
Konfigurasjonen leses bare under trening. Utenfor trening, og valideringen Lightning kjører før første steg er utenfor trening, brukes en miljøvariabel som står på én. Å sette den løste valideringen:
export ANEMOI_INFERENCE_NUM_CHUNKS=16
export ANEMOI_INFERENCE_NUM_CHUNKS_MAPPER=32
Med det kom to medlemmer gjennom valideringen og inn i sitt første treningssteg, som første gang noensinne. Der manglet det under én gigabyte.
Å flytte data til vertsminnet. anemoi har et valg for å legge mellomresultater i vanlig minne i stedet for på kortet. For koderen og dekoderen er det valget ødelagt i denne versjonen, fordi det prøver å pakke inn lagene før de er bygget. For prosessoren virket det, men forsøket ble tretten ganger tregere og sparte bare rundt 2 GB, siden prosessoren allerede regner ut det den trenger på nytt i stedet for å lagre det.
Mer omregning. Jeg pakket inn dekoderen og oppmerksomheten i den slik at mellomresultater skulle regnes på nytt i bakoverpasseringen i stedet for å lagres. Toppen steg fra 134 til 141 GB. anemoi gjør allerede dette for hele koderen og dekoderen, så jeg la bare mer arbeid oppå.
| endring | resultat |
|---|---|
| seksten biter i konfigurasjonen | samme allokering, til byten |
| biter satt via miljøvariabel | to medlemmer gjennom validering, trening manglet under 1 GB |
| mellomresultater i vertsminnet | tretten ganger tregere, 2 GB spart |
| ekstra omregning i dekoderen | toppen steg fra 134 til 141 GB |
| medlemmene ett om gangen | tre medlemmer trener |
Tallet som løste det
Allokeringen som feilet til slutt var 10,37 GiB. Den var like stor med seksten, trettito og sekstifire biter. Et tall som ikke rikker seg når man endrer det man tror styrer det, er et tall som handler om noe annet.
Så jeg regnet på det. Én tensor med 1024 kanaler i float32 over dekoderens punkter, for to stablede medlemmer:
2 × 1 359 281 punkter × 1024 kanaler × 4 byte = 11 141 120 000 byte = 10,37 GiB
Treff på byten. Det var ikke en tensor over kantene i nettet, som bitene deler opp, men en over punktene. Hver bit samler resultatet sitt inn i en tensor over alle punktene, uansett hvor få kanter den selv har. Ingen oppdeling av kantene kunne noen gang gjøre den mindre.
Det var dette tallet som hadde blitt stående gjennom alle forsøkene over. Ingen av dem hadde angrepet det faktiske problemet: at medlemmene var stablet.
Løsningen
Kjør medlemmene gjennom modellen ett om gangen, og sett resultatene sammen igjen før tapet regnes ut:
members = x.shape[2]
outputs = [
original(self, x[:, :, i:i + 1], fcstep=fcstep, **kwargs)
for i in range(members)
]
return torch.cat(outputs, dim=1)
Det halverer hver tensor over punktene ved to medlemmer, og mer ved flere. Matematikken er den samme. Lagnormaliseringen er per prøve, støyen trekkes per prøve, og ingenting i foroverpasseringen kobler ett medlem til et annet. Tapet ser fortsatt alle medlemmene sammen, så ensembleskåren er uendret. Det eneste som skiller seg er hvilke tilfeldige tall støyen får, ikke fordelingen de trekkes fra.
Endringen ligger i en liten fil i prosjektet og brukes likt i begge armene, slik at sammenlikningen mellom dem ikke påvirkes.
Hva det kostet
| halearm, 60 steg | kontrollarm, 5 steg | |
|---|---|---|
| høyeste minnebruk | 142 987 MiB | 143 027 MiB |
| tid per steg | 33 s | 35 s |
| treningstap | 1,85 | 1,92 |
| valideringstap | 2,29 | 2,26 |
Ensemblet har tre medlemmer, ikke fire som opplegget var tegnet for. Fire går fortsatt tom, nå i dekoderens foroverlag. Begge armene bruker tre, så sammenlikningen holder, men ensembleskåren blir svakere, og det skal stå i oppgaven.
Minnemarginen er under én gigabyte. Den holdt seg nøyaktig lik over seksti steg, så den kryper ikke, men det er ingen luft å gå på.
Og hver arm tar rundt 31 timer, inkludert én gjennomgang av valideringsperioden.
Det jeg tar med meg
Tre av fem forsøk var rimelige ideer som angrep feil ting. Det som skilte det som virket fra det som ikke gjorde det, var ikke flere forsøk, men å ta allokeringen som feilet på alvor som et tall, og regne ut hvilken tensor som har akkurat den størrelsen.
