@@ -126,14 +126,21 @@ def select(W, N, hess=None, metric="mse"):
126126 return torch .stack (cands )[ch , torch .arange (W .shape [0 ], device = W .device )]
127127
128128# ── модель (параметризовано env: SEED/STEPS — трек A, NL/DMODEL — трек B) ──
129- device = 'cuda' ; VOCAB = 1024 ; SEQ = 1024 ; BATCH = 48
129+ # АНТИ-OOM: BATCH/SEQ параметризованы + авто-снижение для крупной модели (урок: NL=12/d=768 выбил 31ГБ).
130+ os .environ .setdefault ("PYTORCH_CUDA_ALLOC_CONF" , "expandable_segments:True" )
131+ device = 'cuda' ; VOCAB = 1024
130132SEED = int (os .environ .get ("SEED" , 42 ))
131133D = int (os .environ .get ("DMODEL" , 512 ))
132134NL = int (os .environ .get ("NL" , 9 ))
133135STEPS = int (os .environ .get ("STEPS" , 3000 ))
134136assert D % 8 == 0 , "DMODEL должен быть кратен nhead=8"
137+ # авто-бюджет памяти: целевой ~ NL*D*BATCH*SEQ активаций; держим под 512*9*48*1024
138+ _budget = 9 * 512 * 48 * 1024
139+ _auto_bs = max (8 , min (48 , _budget // (NL * D * 1024 )))
140+ BATCH = int (os .environ .get ("BATCH" , _auto_bs ))
141+ SEQ = int (os .environ .get ("SEQ" , 1024 ))
135142RUN_TAG = f"seed{ SEED } _nl{ NL } _d{ D } _st{ STEPS } "
136- print (f"[config] { RUN_TAG } (трек A=сид/чекпоинт, трек B=NL/DMODEL)" , flush = True )
143+ print (f"[config] { RUN_TAG } BATCH= { BATCH } SEQ= { SEQ } (трек A=сид/чекпоинт, трек B=NL/DMODEL)" , flush = True )
137144os .chdir ("/workspace" )
138145if not os .path .exists ("parameter-golf" ):
139146 os .system ("git clone --depth 1 https://github.com/openai/parameter-golf.git" )
@@ -167,14 +174,15 @@ def forward(s,x):
167174 return s .h (s .f (h ))
168175
169176print (f"GPU: { torch .cuda .get_device_name (0 )} | torch { torch .__version__ } " )
170- print (f"Training { NL } L d={ D } { STEPS } steps..." )
177+ print (f"Training { NL } L d={ D } { STEPS } steps (BATCH= { BATCH } ) ..." )
171178torch .manual_seed (SEED ); model = Model ().to (device )
179+ _last_loss = None
172180op = torch .optim .AdamW (model .parameters (),lr = 0.003 ,weight_decay = 0.1 ,betas = (0.95 ,0.95 ))
173181for s in range (STEPS + 1 ):
174182 idx = torch .randint (0 ,len (train_t )- SEQ - 1 ,(BATCH ,))
175183 x = torch .stack ([train_t [i :i + SEQ ] for i in idx ]).to (device )
176184 y = torch .stack ([train_t [i + 1 :i + SEQ + 1 ] for i in idx ]).to (device )
177- loss = F .cross_entropy (model (x ).reshape (- 1 ,VOCAB ),y .reshape (- 1 ))
185+ loss = F .cross_entropy (model (x ).reshape (- 1 ,VOCAB ),y .reshape (- 1 )); _last_loss = float ( loss )
178186 op .zero_grad (); loss .backward (); torch .nn .utils .clip_grad_norm_ (model .parameters (),1.0 ); op .step ()
179187 if s % 1000 == 0 : print (f" { s } /{ STEPS } : loss={ loss .item ():.4f} " , flush = True )
180188model .eval ()
@@ -265,6 +273,16 @@ def deep_ffn(n):
265273 seed = SEED , run_tag = RUN_TAG )}
266274
267275restore_fp32 (); report ["fp32" ] = eval_bpb ("FP32 baseline" )
276+ report ["meta" ]["train_last_loss" ] = _last_loss
277+ report ["meta" ]["baseline_bpt" ] = report ["fp32" ]["bits_per_token" ]
278+ # GUARD (урок seed=123): коллапс в память (loss→0, baseline BPT→0) делает замер НЕВАЛИДНЫМ:
279+ # квантование нечего портить на вырожденной модели → все ΔBPT≈0 артефактно, не вывод.
280+ _valid = report ["fp32" ]["bits_per_token" ] >= 1.0
281+ report ["meta" ]["baseline_valid" ] = bool (_valid )
282+ if not _valid :
283+ print (f"\n ⚠ BASELINE НЕВАЛИДЕН: FP32 BPT={ report ['fp32' ]['bits_per_token' ]:.5f} < 1.0 "
284+ f"(train loss={ _last_loss :.4f} ). Модель сколлапсировала в память на этом сиде/STEPS."
285+ f"\n Замер ΔBPT НЕВАЛИДЕН (квантовать нечего). Сменить SEED или уменьшить STEPS." , flush = True )
268286
269287for N in (4 , 6 , 8 ):
270288 print (f"\n --- { N } -bit ---" )
0 commit comments