w-ahmad commited on
Commit
ceda489
·
verified ·
1 Parent(s): cfe0147

Auto upload zain 2026-08-14T23:11:25.111342

Browse files
zain/Activation/README.md ADDED
@@ -0,0 +1 @@
 
 
1
+ # Activation
zain/Activation/__pycache__/exp.cpython-311.pyc ADDED
Binary file (62.8 kB). View file
 
zain/Activation/exp.py ADDED
@@ -0,0 +1,921 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # =====================================================================
2
+ # exp.py – FULL FILE, ALL FIXES INCLUDED (hook requires_grad check)
3
+ # =====================================================================
4
+
5
+ import math
6
+ import os
7
+ import time
8
+ import json
9
+ import re
10
+ from pathlib import Path
11
+ from itertools import chain
12
+ from typing import Dict, Callable, Optional, List, Any, Tuple
13
+
14
+ import torch
15
+ import torch.nn as nn
16
+ from transformers import (
17
+ LlamaConfig,
18
+ LlamaPreTrainedModel,
19
+ Trainer,
20
+ TrainerCallback,
21
+ TrainingArguments,
22
+ DataCollatorForLanguageModeling,
23
+ AutoTokenizer,
24
+ set_seed,
25
+ )
26
+ from transformers.models.llama.modeling_llama import (
27
+ LlamaAttention,
28
+ LlamaRMSNorm,
29
+ LlamaRotaryEmbedding,
30
+ )
31
+ from transformers.modeling_outputs import CausalLMOutputWithPast
32
+ from datasets import load_dataset
33
+ from huggingface_hub import snapshot_download
34
+
35
+
36
+ # =============================================================================
37
+ # 1. ACTIVATION REGISTRY
38
+ # =============================================================================
39
+
40
+ class GLUActivationRegistry:
41
+ _registry: Dict[str, Callable[[torch.Tensor], torch.Tensor]] = {}
42
+
43
+ @classmethod
44
+ def register(cls, name: str, fn: Callable[[torch.Tensor], torch.Tensor]) -> None:
45
+ cls._registry[name] = fn
46
+
47
+ @classmethod
48
+ def get(cls, name: str) -> Callable[[torch.Tensor], torch.Tensor]:
49
+ if name not in cls._registry:
50
+ raise KeyError(
51
+ f"Activation '{name}' not found. Available: {list(cls._registry.keys())}"
52
+ )
53
+ return cls._registry[name]
54
+
55
+ # Built-ins
56
+ GLUActivationRegistry.register("silu", nn.functional.silu)
57
+ GLUActivationRegistry.register("swish", nn.functional.silu)
58
+ GLUActivationRegistry.register("relu", nn.functional.relu)
59
+ GLUActivationRegistry.register("gelu", nn.functional.gelu)
60
+ GLUActivationRegistry.register("sigmoid", torch.sigmoid)
61
+ GLUActivationRegistry.register("tanh", torch.tanh)
62
+ GLUActivationRegistry.register("softplus", nn.functional.softplus)
63
+ GLUActivationRegistry.register("linear", lambda x: x)
64
+ GLUActivationRegistry.register("s10", lambda x: x * x * torch.sigmoid(x))
65
+ GLUActivationRegistry.register("w1a", lambda x: x * torch.tanh(x))
66
+
67
+
68
+ # =============================================================================
69
+ # 2. CONFIG
70
+ # =============================================================================
71
+
72
+ class TinyLlamaConfig(LlamaConfig):
73
+ model_type = "tiny_llama"
74
+
75
+ def __init__(
76
+ self,
77
+ mlp_type: str = "glu",
78
+ activation: str = "silu",
79
+ waleed_beta: float = 10.0,
80
+ **kwargs
81
+ ):
82
+ super().__init__(**kwargs)
83
+ self.mlp_type = mlp_type
84
+ self.activation = activation
85
+ self.waleed_beta = waleed_beta
86
+ if self.num_key_value_heads != self.num_attention_heads:
87
+ raise ValueError(
88
+ f"Pure MHA required: num_key_value_heads ({self.num_key_value_heads}) "
89
+ f"must equal num_attention_heads ({self.num_attention_heads})."
90
+ )
91
+
92
+
93
+ # =============================================================================
94
+ # 3. MODEL (with override hook)
95
+ # =============================================================================
96
+
97
+ class TinyLlamaMLP(nn.Module):
98
+ override_active = False
99
+ override_value = -100.0
100
+
101
+ def __init__(self, config: TinyLlamaConfig):
102
+ super().__init__()
103
+ self.hidden_size = config.hidden_size
104
+ self.intermediate_size = config.intermediate_size
105
+ self.mlp_type = config.mlp_type
106
+ self.activation_name = config.activation
107
+ self.waleed_beta = getattr(config, "waleed_beta", 10.0)
108
+
109
+ if self.mlp_type == "glu":
110
+ effective_intermediate = self.intermediate_size
111
+ elif self.mlp_type == "mlp":
112
+ effective_intermediate = int(self.intermediate_size * 1.5)
113
+ print(f"[MLP] Auto‑scaled intermediate_size from {self.intermediate_size} to {effective_intermediate}")
114
+ else:
115
+ raise ValueError(f"Unknown mlp_type: {self.mlp_type}")
116
+
117
+ self.effective_intermediate = effective_intermediate
118
+
119
+ if self.mlp_type == "glu":
120
+ self.gate_proj = nn.Linear(self.hidden_size, effective_intermediate, bias=False)
121
+ self.up_proj = nn.Linear(self.hidden_size, effective_intermediate, bias=False)
122
+ else:
123
+ self.up_proj = nn.Linear(self.hidden_size, effective_intermediate, bias=False)
124
+
125
+ self.down_proj = nn.Linear(effective_intermediate, self.hidden_size, bias=False)
126
+
127
+ # Activation handling
128
+ if self.mlp_type == "glu":
129
+ if self.activation_name in ("situglu", "waleed", "situglu_low", "waleedglu_low"):
130
+ self.act_fn = None
131
+ elif self.activation_name == "waleed10":
132
+ self.act_fn = GLUActivationRegistry.get("linear")
133
+ elif self.activation_name == "silu-waleed10":
134
+ self.act_fn = GLUActivationRegistry.get("silu")
135
+ else:
136
+ self.act_fn = GLUActivationRegistry.get(self.activation_name)
137
+ else:
138
+ if self.activation_name in ("situglu", "waleed", "situglu_low", "waleedglu_low"):
139
+ raise ValueError(f"Activation '{self.activation_name}' requires GLU.")
140
+ elif self.activation_name == "waleed10":
141
+ self.act_fn = GLUActivationRegistry.get("linear")
142
+ elif self.activation_name == "silu-waleed10":
143
+ self.act_fn = GLUActivationRegistry.get("silu")
144
+ else:
145
+ self.act_fn = GLUActivationRegistry.get(self.activation_name)
146
+
147
+ if self.activation_name in ("situglu_low", "waleedglu_low"):
148
+ self.beta1 = 2.5
149
+ self.beta2 = 4.0
150
+ else:
151
+ self.beta1 = 4.0
152
+ self.beta2 = 25.0
153
+
154
+ self.is_situglu = self.activation_name in ("situglu", "situglu_low")
155
+ self.is_waleed = self.activation_name in ("waleed", "waleedglu_low")
156
+ self.is_waleed10 = self.activation_name in ("waleed10", "silu-waleed10")
157
+ self.has_sigmoid_gate = self.activation_name.startswith("situglu")
158
+
159
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
160
+ if self.mlp_type == "glu":
161
+ gate = self.gate_proj(x)
162
+ up = self.up_proj(x)
163
+
164
+ if self.is_situglu or self.is_waleed:
165
+ if self.has_sigmoid_gate:
166
+ gate = self.beta1 * torch.tanh(gate / self.beta1) * torch.sigmoid(gate)
167
+ else:
168
+ gate = self.beta1 * torch.tanh(gate / self.beta1)
169
+ up = self.beta2 * torch.tanh(up / self.beta2)
170
+ hidden = gate * up
171
+ else:
172
+ hidden = self.act_fn(gate) * up
173
+
174
+ out = self.down_proj(hidden)
175
+
176
+ else: # mlp
177
+ hidden = self.act_fn(self.up_proj(x))
178
+ out = self.down_proj(hidden)
179
+
180
+ if self.is_waleed10:
181
+ out = self.waleed_beta * torch.tanh(out / self.waleed_beta)
182
+
183
+ # FIX: only register hook if the output tensor requires gradients
184
+ if TinyLlamaMLP.override_active and out.requires_grad:
185
+ out.register_hook(lambda grad: TinyLlamaMLP.override_value * torch.ones_like(grad))
186
+
187
+ return out
188
+
189
+
190
+ # ----------------------------------------------------------------------------
191
+ # DECODER LAYER, ATTENTION MASK, MODEL
192
+ # ----------------------------------------------------------------------------
193
+
194
+ class TinyLlamaDecoderLayer(nn.Module):
195
+ def __init__(self, config: TinyLlamaConfig, layer_idx: int):
196
+ super().__init__()
197
+ self.hidden_size = config.hidden_size
198
+ self.self_attn = LlamaAttention(config=config, layer_idx=layer_idx)
199
+ self.mlp = TinyLlamaMLP(config)
200
+ self.input_layernorm = LlamaRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
201
+ self.post_attention_layernorm = LlamaRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
202
+ self.residual_pre_attn = nn.Identity()
203
+ self.residual_post_attn = nn.Identity()
204
+ self.residual_post_mlp = nn.Identity()
205
+
206
+ def forward(self, hidden_states, attention_mask=None, position_ids=None, position_embeddings=None, **kwargs):
207
+ residual = hidden_states
208
+ hidden_states = self.residual_pre_attn(hidden_states)
209
+ hidden_states = self.input_layernorm(hidden_states)
210
+ attn_out = self.self_attn(
211
+ hidden_states=hidden_states,
212
+ attention_mask=attention_mask,
213
+ position_ids=position_ids,
214
+ position_embeddings=position_embeddings,
215
+ )[0]
216
+ hidden_states = residual + attn_out
217
+ hidden_states = self.residual_post_attn(hidden_states)
218
+
219
+ residual = hidden_states
220
+ hidden_states = self.post_attention_layernorm(hidden_states)
221
+ hidden_states = self.mlp(hidden_states)
222
+ hidden_states = residual + hidden_states
223
+ hidden_states = self.residual_post_mlp(hidden_states)
224
+ return (hidden_states,)
225
+
226
+
227
+ def _build_causal_mask(attention_mask, seq_len, dtype, device):
228
+ min_value = torch.finfo(dtype).min
229
+ causal = torch.full((seq_len, seq_len), fill_value=min_value, dtype=dtype, device=device)
230
+ causal = torch.triu(causal, diagonal=1)
231
+ causal = causal[None, None, :, :]
232
+ if attention_mask is None:
233
+ batch_size = 1
234
+ return causal.expand(batch_size, 1, seq_len, seq_len)
235
+ batch_size = attention_mask.shape[0]
236
+ causal = causal.expand(batch_size, 1, seq_len, seq_len).clone()
237
+ padding = attention_mask[:, None, None, :].to(device) == 0
238
+ causal = causal.masked_fill(padding, min_value)
239
+ return causal
240
+
241
+
242
+ _MASK_PRINTED = False
243
+
244
+
245
+ class TinyLlamaModel(LlamaPreTrainedModel):
246
+ config_class = TinyLlamaConfig
247
+
248
+ def __init__(self, config: TinyLlamaConfig):
249
+ super().__init__(config)
250
+ self.padding_idx = config.pad_token_id
251
+ self.vocab_size = config.vocab_size
252
+ self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
253
+ self.layers = nn.ModuleList([TinyLlamaDecoderLayer(config, i) for i in range(config.num_hidden_layers)])
254
+ self.norm = LlamaRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
255
+ self.rotary_emb = LlamaRotaryEmbedding(config=config)
256
+ self.post_init()
257
+
258
+ def forward(self, input_ids=None, attention_mask=None, position_ids=None, inputs_embeds=None, return_dict=None, **kwargs):
259
+ global _MASK_PRINTED
260
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
261
+ if inputs_embeds is None:
262
+ inputs_embeds = self.embed_tokens(input_ids)
263
+ if position_ids is None:
264
+ seq_len = inputs_embeds.shape[1]
265
+ position_ids = torch.arange(seq_len, device=inputs_embeds.device).unsqueeze(0).expand(inputs_embeds.shape[0], -1)
266
+ hidden_states = inputs_embeds
267
+ position_embeddings = self.rotary_emb(hidden_states, position_ids)
268
+ seq_len = hidden_states.shape[1]
269
+ causal_mask = _build_causal_mask(attention_mask, seq_len, hidden_states.dtype, hidden_states.device)
270
+ if not _MASK_PRINTED:
271
+ print("[INFO] Causal mask (float with -inf) applied to all attention layers.")
272
+ _MASK_PRINTED = True
273
+ for decoder_layer in self.layers:
274
+ layer_outputs = decoder_layer(
275
+ hidden_states,
276
+ attention_mask=causal_mask,
277
+ position_ids=position_ids,
278
+ position_embeddings=position_embeddings,
279
+ )
280
+ hidden_states = layer_outputs[0]
281
+ hidden_states = self.norm(hidden_states)
282
+ if not return_dict:
283
+ return (hidden_states,)
284
+ return {"last_hidden_state": hidden_states, "hidden_states": None, "attentions": None}
285
+
286
+
287
+ class TinyLlamaForCausalLM(LlamaPreTrainedModel):
288
+ config_class = TinyLlamaConfig
289
+ _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"}
290
+
291
+ def __init__(self, config: TinyLlamaConfig):
292
+ super().__init__(config)
293
+ self.model = TinyLlamaModel(config)
294
+ self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
295
+ if config.tie_word_embeddings:
296
+ self.lm_head.weight = self.model.embed_tokens.weight
297
+ self.post_init()
298
+
299
+ def get_input_embeddings(self):
300
+ return self.model.embed_tokens
301
+
302
+ def set_input_embeddings(self, value):
303
+ self.model.embed_tokens = value
304
+
305
+ def get_output_embeddings(self):
306
+ return self.lm_head
307
+
308
+ def forward(self, input_ids=None, attention_mask=None, position_ids=None, inputs_embeds=None,
309
+ labels=None, return_dict=None, **kwargs):
310
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
311
+ outputs = self.model(input_ids=input_ids, attention_mask=attention_mask, position_ids=position_ids,
312
+ inputs_embeds=inputs_embeds, return_dict=return_dict)
313
+ hidden_states = outputs["last_hidden_state"] if return_dict else outputs[0]
314
+ logits = self.lm_head(hidden_states)
315
+ loss = None
316
+ if labels is not None:
317
+ shift_logits = logits[..., :-1, :].contiguous()
318
+ shift_labels = labels[..., 1:].contiguous()
319
+ loss_fct = nn.CrossEntropyLoss()
320
+ loss = loss_fct(shift_logits.view(-1, self.config.vocab_size), shift_labels.view(-1))
321
+ if not return_dict:
322
+ output = (logits,) + outputs[1:]
323
+ return (loss,) + output if loss is not None else output
324
+ return CausalLMOutputWithPast(loss=loss, logits=logits, past_key_values=None, hidden_states=None, attentions=None)
325
+
326
+ def prepare_inputs_for_generation(self, input_ids, past_key_values=None, attention_mask=None, **kwargs):
327
+ if past_key_values:
328
+ input_ids = input_ids[:, -1:]
329
+ position_ids = kwargs.get("position_ids")
330
+ if attention_mask is not None and position_ids is None:
331
+ position_ids = attention_mask.long().cumsum(-1) - 1
332
+ position_ids.masked_fill_(attention_mask == 0, 1)
333
+ if past_key_values:
334
+ position_ids = position_ids[:, -1].unsqueeze(-1)
335
+ return {
336
+ "input_ids": input_ids,
337
+ "position_ids": position_ids,
338
+ "past_key_values": past_key_values,
339
+ "attention_mask": attention_mask,
340
+ }
341
+
342
+
343
+ # =============================================================================
344
+ # 4. HF CHECKPOINT FETCHER
345
+ # =============================================================================
346
+
347
+ def fetch_latest_checkpoint_from_hub(
348
+ repo_id: str,
349
+ subpath: str,
350
+ variant: str,
351
+ checkpoint_step: Optional[int] = None
352
+ ) -> str:
353
+ remote_prefix = f"{subpath}/{variant}_run" if subpath else f"{variant}_run"
354
+ if checkpoint_step is not None:
355
+ full_remote_path = f"{remote_prefix}/checkpoint-{checkpoint_step}"
356
+ print(f"[Hub] Fetching specific checkpoint: {repo_id}/{full_remote_path}")
357
+ local_root = snapshot_download(
358
+ repo_id=repo_id,
359
+ allow_patterns=[f"{full_remote_path}/*"],
360
+ local_dir_use_symlinks=False,
361
+ )
362
+ checkpoint_local_path = os.path.join(local_root, full_remote_path)
363
+ if not os.path.exists(checkpoint_local_path):
364
+ raise RuntimeError(f"Downloaded checkpoint not found at {checkpoint_local_path}")
365
+ return checkpoint_local_path
366
+
367
+ print(f"[Hub] Downloading entire run folder: {repo_id}/{remote_prefix}")
368
+ local_root = snapshot_download(
369
+ repo_id=repo_id,
370
+ allow_patterns=[f"{remote_prefix}/*"],
371
+ local_dir_use_symlinks=False,
372
+ )
373
+ run_local_path = os.path.join(local_root, remote_prefix)
374
+ if not os.path.exists(run_local_path):
375
+ raise RuntimeError(f"Run folder not found at {run_local_path}")
376
+
377
+ checkpoints = []
378
+ for item in os.listdir(run_local_path):
379
+ if item.startswith("checkpoint-") and os.path.isdir(os.path.join(run_local_path, item)):
380
+ match = re.match(r"checkpoint-(\d+)", item)
381
+ if match:
382
+ step = int(match.group(1))
383
+ checkpoints.append((step, item))
384
+
385
+ if not checkpoints:
386
+ raise RuntimeError(f"No checkpoint folders found in {run_local_path}")
387
+
388
+ latest_step, latest_name = max(checkpoints, key=lambda x: x[0])
389
+ print(f"[Hub] Latest checkpoint found: step {latest_step}")
390
+ return os.path.join(run_local_path, latest_name)
391
+
392
+
393
+ # =============================================================================
394
+ # 5. RESUME + FREEZE + OVERRIDE CALLBACK (FIXED)
395
+ # =============================================================================
396
+
397
+ class ResumeFreezeOverrideCallback(TrainerCallback):
398
+ def __init__(self, override_value: Optional[float] = None):
399
+ self.override_value = override_value
400
+ self._trainer = None # Will be set manually
401
+
402
+ def on_train_begin(self, args, state, control, **kwargs):
403
+ # Try multiple ways to get the trainer
404
+ trainer = kwargs.get('trainer')
405
+ if trainer is None:
406
+ trainer = getattr(self, '_trainer', None)
407
+ if trainer is None:
408
+ trainer = getattr(self, 'trainer', None)
409
+ if trainer is None:
410
+ raise ValueError("Trainer not accessible in callback")
411
+
412
+ model = trainer.model
413
+
414
+ # ----- FREEZE ALL EXCEPT MLP PROJECTIONS -----
415
+ for name, param in model.named_parameters():
416
+ if any(x in name for x in ["gate_proj", "up_proj", "down_proj"]):
417
+ param.requires_grad = True
418
+ else:
419
+ param.requires_grad = False
420
+
421
+ trainable_params = [p for p in model.parameters() if p.requires_grad]
422
+
423
+ # ----- REBUILD OPTIMIZER -----
424
+ from torch.optim import AdamW
425
+ new_optimizer = AdamW(
426
+ trainable_params,
427
+ lr=args.learning_rate,
428
+ weight_decay=args.weight_decay,
429
+ betas=(args.adam_beta1, args.adam_beta2),
430
+ eps=args.adam_epsilon,
431
+ )
432
+ trainer.optimizer = new_optimizer
433
+
434
+ # ----- KEEP EXISTING SCHEDULER (re‑attach) -----
435
+ if trainer.lr_scheduler is not None:
436
+ trainer.lr_scheduler.optimizer = new_optimizer
437
+
438
+ # ----- ACTIVATE OVERRIDE -----
439
+ if self.override_value is not None:
440
+ TinyLlamaMLP.override_active = True
441
+ TinyLlamaMLP.override_value = self.override_value
442
+ print(f"[Override] Activated with value = {self.override_value}")
443
+ else:
444
+ print("[Freeze] MLP projections frozen; no gradient override.")
445
+
446
+
447
+ # =============================================================================
448
+ # 6. MONITORING
449
+ # =============================================================================
450
+
451
+ class StatsEngine:
452
+ @staticmethod
453
+ def compute(tensor: torch.Tensor, user_limit: float, dtype_ratio: float) -> Dict[str, float]:
454
+ with torch.no_grad():
455
+ abs_t = tensor.abs()
456
+ dtype_info = torch.finfo(tensor.dtype)
457
+ dtype_limit = dtype_ratio * dtype_info.max if not torch.isinf(torch.tensor(dtype_info.max)) else float("inf")
458
+ return {
459
+ "norm": tensor.norm(2).item(),
460
+ "mean": tensor.mean().item(),
461
+ "std": tensor.std().item(),
462
+ "max_abs": abs_t.max().item(),
463
+ "frac_near_dtype_limit": (abs_t > dtype_limit).float().mean().item() if not math.isinf(dtype_limit) else 0.0,
464
+ "frac_near_user_limit": (abs_t > user_limit).float().mean().item(),
465
+ "min": tensor.min().item(),
466
+ "max": tensor.max().item(),
467
+ "range": tensor.max().item() - tensor.min().item(),
468
+ }
469
+
470
+
471
+ class StepAccumulator:
472
+ def __init__(self):
473
+ self.tensors: Dict[str, Dict[str, float]] = {}
474
+
475
+ def add(self, name: str, numel: int, stats: Dict[str, float]):
476
+ new_entry = {"numel": numel, **stats}
477
+ existing = self.tensors.get(name)
478
+ self.tensors[name] = new_entry if existing is None else self._merge_entry(existing, new_entry)
479
+
480
+ @staticmethod
481
+ def _merge_entry(a: Dict[str, float], b: Dict[str, float]) -> Dict[str, float]:
482
+ total_n = a["numel"] + b["numel"]
483
+ if total_n == 0:
484
+ return a
485
+ norm = math.sqrt(a["norm"] ** 2 + b["norm"] ** 2)
486
+ max_abs = max(a["max_abs"], b["max_abs"])
487
+ mean = (a["mean"] * a["numel"] + b["mean"] * b["numel"]) / total_n
488
+ ex2 = (a["numel"] * (a["std"] ** 2 + a["mean"] ** 2) + b["numel"] * (b["std"] ** 2 + b["mean"] ** 2)) / total_n
489
+ std = math.sqrt(max(0.0, ex2 - mean ** 2))
490
+ frac_dtype = (a["frac_near_dtype_limit"] * a["numel"] + b["frac_near_dtype_limit"] * b["numel"]) / total_n
491
+ frac_user = (a["frac_near_user_limit"] * a["numel"] + b["frac_near_user_limit"] * b["numel"]) / total_n
492
+ t_min = min(a.get("min", float("inf")), b.get("min", float("inf")))
493
+ t_max = max(a.get("max", float("-inf")), b.get("max", float("-inf")))
494
+ return {
495
+ "numel": total_n,
496
+ "norm": norm,
497
+ "mean": mean,
498
+ "std": std,
499
+ "max_abs": max_abs,
500
+ "frac_near_dtype_limit": frac_dtype,
501
+ "frac_near_user_limit": frac_user,
502
+ "min": t_min,
503
+ "max": t_max,
504
+ "range": t_max - t_min,
505
+ }
506
+
507
+ def clear(self):
508
+ self.tensors.clear()
509
+
510
+ def _aggregate(self, entries: Dict[str, Dict[str, float]]) -> Dict[str, float]:
511
+ if not entries:
512
+ return {}
513
+ numels = [e["numel"] for e in entries.values()]
514
+ total_n = sum(numels)
515
+ norm = math.sqrt(sum(e["norm"] ** 2 for e in entries.values()))
516
+ max_abs = max(e["max_abs"] for e in entries.values())
517
+ mean = sum(e["mean"] * e["numel"] for e in entries.values()) / total_n
518
+ ex2 = sum(e["numel"] * (e["std"] ** 2 + e["mean"] ** 2) for e in entries.values()) / total_n
519
+ std = math.sqrt(max(0.0, ex2 - mean ** 2))
520
+ frac_dtype = sum(e["frac_near_dtype_limit"] * e["numel"] for e in entries.values()) / total_n
521
+ frac_user = sum(e["frac_near_user_limit"] * e["numel"] for e in entries.values()) / total_n
522
+ t_min = min(e.get("min", float("inf")) for e in entries.values())
523
+ t_max = max(e.get("max", float("-inf")) for e in entries.values())
524
+ return {
525
+ "norm": norm,
526
+ "mean": mean,
527
+ "std": std,
528
+ "max_abs": max_abs,
529
+ "frac_near_dtype_limit": frac_dtype,
530
+ "frac_near_user_limit": frac_user,
531
+ "min": t_min,
532
+ "max": t_max,
533
+ "range": t_max - t_min,
534
+ }
535
+
536
+ def get_global_stats(self) -> Dict[str, float]:
537
+ return self._aggregate(self.tensors)
538
+
539
+ def get_layer_stats(self, layer_prefix: str) -> Dict[str, float]:
540
+ entries = {k: v for k, v in self.tensors.items() if k.startswith(layer_prefix + ".")}
541
+ return self._aggregate(entries)
542
+
543
+
544
+ class HookRegistry:
545
+ def __init__(self, model: nn.Module):
546
+ self.model = model
547
+ self.handles = []
548
+ self.active = False
549
+
550
+ def attach_forward(self, module_patterns, accumulator, user_limit, dtype_ratio):
551
+ for name, module in self.model.named_modules():
552
+ if any(re.search(p, name) for p in module_patterns):
553
+ h = module.register_forward_hook(
554
+ self._make_forward_hook(name, accumulator, user_limit, dtype_ratio)
555
+ )
556
+ self.handles.append(h)
557
+
558
+ def attach_backward(self, param_patterns, accumulator, user_limit, dtype_ratio):
559
+ for name, param in self.model.named_parameters():
560
+ if not param.requires_grad:
561
+ continue
562
+ if param_patterns and not any(re.search(p, name) for p in param_patterns):
563
+ continue
564
+ h = param.register_hook(
565
+ self._make_backward_hook(f"grad.{name}", accumulator, user_limit, dtype_ratio)
566
+ )
567
+ self.handles.append(h)
568
+
569
+ def _make_forward_hook(self, module_name, accumulator, user_limit, dtype_ratio):
570
+ def hook(module, inp, out):
571
+ if not self.active:
572
+ return
573
+ if isinstance(out, dict):
574
+ out = out.get("last_hidden_state")
575
+ if not torch.is_tensor(out):
576
+ return
577
+ stats = StatsEngine.compute(out.detach(), user_limit, dtype_ratio)
578
+ accumulator.add(f"act.{module_name}", out.numel(), stats)
579
+ return hook
580
+
581
+ def _make_backward_hook(self, param_name, accumulator, user_limit, dtype_ratio):
582
+ def hook(grad):
583
+ if not self.active:
584
+ return
585
+ stats = StatsEngine.compute(grad.detach(), user_limit, dtype_ratio)
586
+ accumulator.add(param_name, grad.numel(), stats)
587
+ return hook
588
+
589
+ def set_active(self, active: bool):
590
+ self.active = active
591
+
592
+ def clear(self):
593
+ for h in self.handles:
594
+ h.remove()
595
+ self.handles.clear()
596
+
597
+
598
+ class StabilityMonitorCallback(TrainerCallback):
599
+ def __init__(
600
+ self,
601
+ model: nn.Module,
602
+ monitor_every_n_steps: int = 10,
603
+ module_patterns: Optional[List[str]] = None,
604
+ param_patterns: Optional[List[str]] = None,
605
+ user_limits: Optional[Dict[str, float]] = None,
606
+ dtype_proximity_ratio: float = 0.9,
607
+ log_scope: Optional[Dict[str, bool]] = None,
608
+ monitor_during_eval: bool = False,
609
+ ):
610
+ self.model = model
611
+ self.monitor_every_n_steps = monitor_every_n_steps
612
+ self.module_patterns = module_patterns or [".*mlp.*", ".*residual.*"]
613
+ self.param_patterns = param_patterns or self.module_patterns
614
+ self.user_limits = user_limits or {"grad": 1.0, "param": 100.0, "act": 50.0}
615
+ self.dtype_ratio = dtype_proximity_ratio
616
+ self.log_scope = log_scope or {"global": True, "per_layer": True, "per_tensor": False}
617
+ self.monitor_during_eval = monitor_during_eval
618
+
619
+ self.accumulator = StepAccumulator()
620
+ self.hooks = HookRegistry(model)
621
+ self.hooks.attach_forward(
622
+ self.module_patterns,
623
+ self.accumulator,
624
+ self.user_limits["act"],
625
+ self.dtype_ratio,
626
+ )
627
+ self.hooks.attach_backward(
628
+ self.param_patterns,
629
+ self.accumulator,
630
+ self.user_limits["grad"],
631
+ self.dtype_ratio,
632
+ )
633
+ self.pending_metrics = None
634
+
635
+ def _should_monitor(self, state):
636
+ return state.global_step % self.monitor_every_n_steps == 0
637
+
638
+ def on_step_begin(self, args, state, control, **kwargs):
639
+ if self._should_monitor(state):
640
+ self.accumulator.clear()
641
+ self.hooks.set_active(True)
642
+
643
+ def on_step_end(self, args, state, control, **kwargs):
644
+ if not self.hooks.active:
645
+ return
646
+ for name, param in self.model.named_parameters():
647
+ if self.param_patterns and not any(re.search(p, name) for p in self.param_patterns):
648
+ continue
649
+ stats = StatsEngine.compute(param.data, self.user_limits["param"], self.dtype_ratio)
650
+ self.accumulator.add(f"param.{name}", param.numel(), stats)
651
+ self.hooks.set_active(False)
652
+ self.pending_metrics = self._build_metrics()
653
+
654
+ def _kind_of(self, name: str) -> str:
655
+ if name.startswith("act."):
656
+ return "act"
657
+ if name.startswith("grad."):
658
+ return "grad"
659
+ if name.startswith("param."):
660
+ return "param"
661
+ return "other"
662
+
663
+ def _strip_kind(self, name: str) -> str:
664
+ if name.startswith("act."):
665
+ return name[4:]
666
+ if name.startswith(("grad.", "param.")):
667
+ return name[5:]
668
+ return name
669
+
670
+ def _build_metrics(self, scope: str = "train") -> Dict[str, float]:
671
+ metrics = {}
672
+ if self.log_scope.get("global", True):
673
+ by_kind = {}
674
+ for k, v in self.accumulator.tensors.items():
675
+ by_kind.setdefault(self._kind_of(k), {})[k] = v
676
+ for kind, entries in by_kind.items():
677
+ stats = self.accumulator._aggregate(entries)
678
+ for kk, vv in stats.items():
679
+ metrics[f"{scope}/global/{kind}/{kk}"] = vv
680
+
681
+ if self.log_scope.get("per_layer", True):
682
+ layer_prefixes = set()
683
+ for name in self.accumulator.tensors:
684
+ clean = self._strip_kind(name)
685
+ parts = clean.split(".")
686
+ for i, p in enumerate(parts):
687
+ if p == "layers" and i + 1 < len(parts):
688
+ prefix = ".".join(parts[: i + 2])
689
+ layer_prefixes.add(prefix)
690
+ for prefix in layer_prefixes:
691
+ by_kind = {}
692
+ for k, v in self.accumulator.tensors.items():
693
+ clean = self._strip_kind(k)
694
+ if clean.startswith(prefix + ".") or clean == prefix:
695
+ by_kind.setdefault(self._kind_of(k), {})[k] = v
696
+ safe = prefix.replace(".", "_")
697
+ for kind, entries in by_kind.items():
698
+ if not entries:
699
+ continue
700
+ stats = self.accumulator._aggregate(entries)
701
+ for kk, vv in stats.items():
702
+ metrics[f"{scope}/layer_{safe}/{kind}/{kk}"] = vv
703
+
704
+ if self.log_scope.get("per_tensor", False):
705
+ for name, stats in self.accumulator.tensors.items():
706
+ safe = name.replace(".", "_")
707
+ for kk, vv in stats.items():
708
+ if kk == "numel":
709
+ continue
710
+ metrics[f"{scope}/tensor_{safe}/{kk}"] = vv
711
+ return metrics
712
+
713
+ def on_log(self, args, state, control, logs=None, **kwargs):
714
+ if logs is not None and self.pending_metrics is not None:
715
+ logs.update(self.pending_metrics)
716
+ try:
717
+ import wandb
718
+ if wandb.run is not None:
719
+ wandb.log(self.pending_metrics, step=state.global_step)
720
+ except ImportError:
721
+ pass
722
+ self.pending_metrics = None
723
+
724
+ def on_prediction_step(self, args, state, control, **kwargs):
725
+ if not self.monitor_during_eval:
726
+ return
727
+ if not self.hooks.active:
728
+ self.accumulator.clear()
729
+ self.hooks.set_active(True)
730
+ for name, param in self.model.named_parameters():
731
+ if self.param_patterns and not any(re.search(p, name) for p in self.param_patterns):
732
+ continue
733
+ stats = StatsEngine.compute(param.data, self.user_limits["param"], self.dtype_ratio)
734
+ self.accumulator.add(f"param.{name}", param.numel(), stats)
735
+ self.pending_metrics = self._build_metrics(scope="eval")
736
+
737
+ def on_evaluate(self, args, state, control, metrics=None, **kwargs):
738
+ self.hooks.set_active(False)
739
+ self.accumulator.clear()
740
+
741
+
742
+ # =============================================================================
743
+ # 7. TIME TRACKER
744
+ # =============================================================================
745
+
746
+ class TimeTrackerCallback(TrainerCallback):
747
+ def __init__(self):
748
+ self.step_start = None
749
+ self.epoch_start = None
750
+ self.total_train_time = 0.0
751
+ self.step_times = []
752
+
753
+ def on_epoch_begin(self, args, state, control, **kwargs):
754
+ self.epoch_start = time.perf_counter()
755
+
756
+ def on_step_begin(self, args, state, control, **kwargs):
757
+ self.step_start = time.perf_counter()
758
+
759
+ def on_step_end(self, args, state, control, **kwargs):
760
+ if self.step_start is not None:
761
+ dt = time.perf_counter() - self.step_start
762
+ self.step_times.append(dt)
763
+ self.total_train_time += dt
764
+ self.step_start = None
765
+
766
+ def on_log(self, args, state, control, logs=None, **kwargs):
767
+ if logs is None:
768
+ return
769
+ logs["train/total_time_seconds"] = self.total_train_time
770
+ if self.step_times:
771
+ recent = self.step_times[-100:]
772
+ logs["train/time_per_step_avg"] = sum(recent) / len(recent)
773
+ if self.epoch_start is not None:
774
+ logs["train/epoch_time_elapsed"] = time.perf_counter() - self.epoch_start
775
+ if state.max_steps and state.global_step > 0:
776
+ avg = self.total_train_time / state.global_step
777
+ remaining = (state.max_steps - state.global_step) * avg
778
+ logs["train/estimated_remaining_minutes"] = remaining / 60.0
779
+
780
+
781
+ # =============================================================================
782
+ # 8. METRICS LOGGER
783
+ # =============================================================================
784
+
785
+ class MetricsLoggerCallback(TrainerCallback):
786
+ def __init__(self, output_dir: str):
787
+ self.output_dir = Path(output_dir)
788
+ self.output_dir.mkdir(parents=True, exist_ok=True)
789
+ self.log_file = self.output_dir / "training_log.jsonl"
790
+
791
+ def on_log(self, args, state, control, logs=None, **kwargs):
792
+ if logs is None:
793
+ return
794
+ entry = {"step": state.global_step, "epoch": state.epoch, "timestamp": time.time(), **logs}
795
+ with open(self.log_file, "a") as f:
796
+ f.write(json.dumps(entry, default=str) + "\n")
797
+
798
+
799
+ # =============================================================================
800
+ # 9. DATA & TRAINER FACTORY
801
+ # =============================================================================
802
+
803
+ def build_dataset(
804
+ tokenizer,
805
+ max_seq_len: int = 512,
806
+ split: str = "train",
807
+ dataset_name: str = "roneneldan/TinyStories",
808
+ max_samples: Optional[int] = None,
809
+ ):
810
+ ds = load_dataset(dataset_name, split=split)
811
+ if max_samples is not None and split == "train":
812
+ ds = ds.select(range(min(max_samples, len(ds))))
813
+
814
+ def tokenize(examples):
815
+ out = tokenizer(examples["text"], add_special_tokens=False)
816
+ eos_id = tokenizer.eos_token_id
817
+ out["input_ids"] = [ids + [eos_id] for ids in out["input_ids"]]
818
+ if "attention_mask" in out:
819
+ out["attention_mask"] = [mask + [1] for mask in out["attention_mask"]]
820
+ return out
821
+
822
+ tokenized = ds.map(tokenize, batched=True, num_proc=4, remove_columns=ds.column_names)
823
+
824
+ def group_texts(examples):
825
+ concatenated = {k: list(chain.from_iterable(examples[k])) for k in examples.keys()}
826
+ total_length = len(concatenated[list(examples.keys())[0]])
827
+ total_length = (total_length // max_seq_len) * max_seq_len
828
+ result = {
829
+ k: [t[i:i + max_seq_len] for i in range(0, total_length, max_seq_len)]
830
+ for k, t in concatenated.items()
831
+ }
832
+ result["labels"] = result["input_ids"].copy()
833
+ return result
834
+
835
+ return tokenized.map(group_texts, batched=True, batch_size=10000, num_proc=4)
836
+
837
+
838
+ def create_trainer(
839
+ model,
840
+ tokenizer,
841
+ config: Dict[str, Any],
842
+ train_dataset,
843
+ eval_dataset=None,
844
+ ):
845
+ tc = config.get("training", {})
846
+ mc = config.get("monitor", {})
847
+ go = config.get("gradient_override", {})
848
+
849
+ wandb_project = tc.get("wandb_project")
850
+ if wandb_project:
851
+ os.environ["WANDB_PROJECT"] = wandb_project
852
+
853
+ args = TrainingArguments(
854
+ output_dir=tc.get("output_dir", "./out"),
855
+ run_name=tc.get("run_name", None),
856
+ num_train_epochs=tc.get("num_train_epochs", 3),
857
+ per_device_train_batch_size=tc.get("per_device_train_batch_size", 16),
858
+ per_device_eval_batch_size=tc.get("per_device_eval_batch_size", 16),
859
+ gradient_accumulation_steps=tc.get("gradient_accumulation_steps", 4),
860
+ learning_rate=tc.get("learning_rate", 3e-4),
861
+ weight_decay=tc.get("weight_decay", 0.0),
862
+ max_grad_norm=tc.get("max_grad_norm", 1.0),
863
+ optim=tc.get("optim", "adamw_torch"),
864
+ warmup_steps=tc.get("warmup_steps", 0),
865
+ lr_scheduler_type=tc.get("lr_scheduler_type", "cosine"),
866
+ bf16=tc.get("bf16", True),
867
+ logging_steps=tc.get("logging_steps", 10),
868
+ eval_strategy=tc.get("eval_strategy", "steps"),
869
+ eval_steps=tc.get("eval_steps", 500),
870
+ save_strategy=tc.get("save_strategy", "steps"),
871
+ save_steps=tc.get("save_steps", 1000),
872
+ load_best_model_at_end=tc.get("load_best_model_at_end", False),
873
+ report_to=tc.get("report_to", "tensorboard"),
874
+ push_to_hub=tc.get("push_to_hub", False),
875
+ hub_model_id=tc.get("hub_model_id", None),
876
+ hub_token=tc.get("hub_token") or os.environ.get("HF_TOKEN"),
877
+ max_steps=tc.get("max_steps", -1),
878
+ seed=tc.get("seed", 42),
879
+ data_seed=tc.get("data_seed", 42),
880
+ remove_unused_columns=False,
881
+ )
882
+
883
+ callbacks = [TimeTrackerCallback()]
884
+
885
+ # ---- FREEZE + OVERRIDE CALLBACK ----
886
+ freeze_mlp = tc.get("freeze_mlp", False)
887
+ override_enabled = go.get("enabled", False)
888
+ if freeze_mlp or override_enabled:
889
+ override_value = go.get("value", -100.0) if override_enabled else None
890
+ callbacks.append(ResumeFreezeOverrideCallback(override_value=override_value))
891
+
892
+ # ---- MONITORING ----
893
+ if mc.get("enabled", True):
894
+ module_patterns = mc.get("module_patterns", [".*mlp.*", ".*residual.*"])
895
+ callbacks.append(
896
+ StabilityMonitorCallback(
897
+ model=model,
898
+ monitor_every_n_steps=mc.get("monitor_every_n_steps", 10),
899
+ module_patterns=module_patterns,
900
+ param_patterns=module_patterns,
901
+ user_limits=mc.get("user_limits", {"grad": 1.0, "param": 100.0, "act": 50.0}),
902
+ dtype_proximity_ratio=mc.get("dtype_proximity_ratio", 0.9),
903
+ log_scope=mc.get("log_scope", {"global": True, "per_layer": True, "per_tensor": False}),
904
+ monitor_during_eval=mc.get("monitor_during_eval", False),
905
+ )
906
+ )
907
+
908
+ callbacks.append(MetricsLoggerCallback(args.output_dir))
909
+
910
+ collator = DataCollatorForLanguageModeling(tokenizer=tokenizer, mlm=False)
911
+
912
+ trainer = Trainer(
913
+ model=model,
914
+ args=args,
915
+ train_dataset=train_dataset,
916
+ eval_dataset=eval_dataset,
917
+ data_collator=collator,
918
+ callbacks=callbacks,
919
+ )
920
+
921
+ return trainer
zain/Activation/llm_analyzer_wandb.py ADDED
@@ -0,0 +1,570 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Multi-LLM Activation & Loss Analyzer with WandB Logging
3
+
4
+ Evaluates multiple language models on WikiText dataset,
5
+ computing per-tensor and global activation statistics (mean, max_abs, std, norm)
6
+ and logging to Weights & Biases with separate runs per model.
7
+
8
+ Requirements:
9
+ pip install transformers torch datasets wandb tqdm
10
+
11
+ Usage:
12
+ export WANDB_PROJECT="llm-activation-analysis"
13
+ export WANDB_API_KEY="your-key"
14
+ python llm_analyzer_wandb.py
15
+ """
16
+
17
+ import math
18
+ import torch
19
+ import torch.nn.functional as F
20
+ from transformers import AutoTokenizer, AutoModelForCausalLM, AutoConfig
21
+ from datasets import load_dataset
22
+ from typing import List, Dict, Optional, Union, Tuple
23
+ from dataclasses import dataclass, asdict
24
+ from collections import defaultdict
25
+ import json
26
+ import warnings
27
+ import os
28
+ from tqdm import tqdm
29
+
30
+ import wandb
31
+
32
+ warnings.filterwarnings("ignore")
33
+
34
+
35
+ @dataclass
36
+ class TensorStats:
37
+ mean: float
38
+ max_abs: float
39
+ std: float
40
+ norm: float
41
+ numel: int
42
+
43
+
44
+ @dataclass
45
+ class ModelResult:
46
+ model_name: str
47
+ loss: float
48
+ perplexity: float
49
+ global_act: TensorStats
50
+ layer_acts: Dict[str, TensorStats]
51
+ num_tokens: int
52
+ num_layers: int
53
+ hidden_size: int
54
+ num_params: int
55
+
56
+
57
+ class ActivationHookManager:
58
+ """Manages forward hooks to capture activations from every tensor."""
59
+
60
+ def __init__(self):
61
+ self.activations = {}
62
+ self.hooks = []
63
+ # Set once per batch via set_attention_mask(); used to exclude
64
+ # padding-token positions from activation statistics.
65
+ self._attention_mask: Optional[torch.Tensor] = None
66
+
67
+ def set_attention_mask(self, attention_mask: Optional[torch.Tensor]):
68
+ """Call once per batch before the forward pass so hooks can mask
69
+ out padding positions when computing stats."""
70
+ self._attention_mask = (
71
+ attention_mask.detach().cpu() if attention_mask is not None else None
72
+ )
73
+
74
+ def _make_hook(self, name: str):
75
+ def hook(module, input, output):
76
+ # Handle different output types
77
+ if isinstance(output, torch.Tensor):
78
+ tensor = output
79
+ elif isinstance(output, tuple) and isinstance(output[0], torch.Tensor):
80
+ tensor = output[0]
81
+ else:
82
+ return
83
+
84
+ # Detach and move to CPU to avoid GPU memory blowup
85
+ self.activations[name] = tensor.detach().cpu().float()
86
+ return hook
87
+
88
+ def register_hooks(self, model: torch.nn.Module):
89
+ """Register hooks on all modules that produce activations."""
90
+ for name, module in model.named_modules():
91
+ # Skip trivial containers
92
+ if len(list(module.children())) == 0 and hasattr(module, 'forward'):
93
+ hook = module.register_forward_hook(self._make_hook(name))
94
+ self.hooks.append(hook)
95
+
96
+ def clear(self):
97
+ self.activations.clear()
98
+
99
+ def remove_hooks(self):
100
+ for hook in self.hooks:
101
+ hook.remove()
102
+ self.hooks.clear()
103
+
104
+ def _select_real_tokens(self, tensor: torch.Tensor) -> torch.Tensor:
105
+ """
106
+ BUG FIX: previously every captured tensor (including activations at
107
+ padding-token positions) was flattened and used as-is. With
108
+ padding_side="left" and a small batch_size, the padding fraction
109
+ varies a lot batch-to-batch, so padding-token activations (which
110
+ are real, non-zero values — not zeros) were silently mixed into
111
+ mean/std/max_abs/norm, biasing exactly the saturation signal this
112
+ script exists to measure.
113
+
114
+ Here we mask out padding positions whenever a tensor's shape is
115
+ consistent with (batch, seq_len, ...) against the stored
116
+ attention_mask (batch, seq_len). Tensors that don't match that
117
+ shape (e.g. a module operating on the pooled/final dimension only)
118
+ are left as-is rather than guessing.
119
+ """
120
+ mask = self._attention_mask
121
+ if mask is None or tensor.dim() < 2:
122
+ return tensor.reshape(-1)
123
+ if tensor.shape[0] != mask.shape[0] or tensor.shape[1] != mask.shape[1]:
124
+ return tensor.reshape(-1)
125
+
126
+ bool_mask = mask.bool()
127
+ # Expand mask across any trailing dims (e.g. hidden_size) and select.
128
+ expand_shape = bool_mask.shape + (1,) * (tensor.dim() - 2)
129
+ bool_mask = bool_mask.view(expand_shape).expand_as(tensor)
130
+ return tensor[bool_mask].reshape(-1)
131
+
132
+ def compute_stats(self) -> Dict[str, TensorStats]:
133
+ """Compute statistics for all captured activations, excluding
134
+ padding-token positions where identifiable."""
135
+ stats = {}
136
+ for name, tensor in self.activations.items():
137
+ if tensor.numel() == 0:
138
+ continue
139
+ flat = self._select_real_tokens(tensor)
140
+ if flat.numel() == 0:
141
+ continue
142
+ stats[name] = TensorStats(
143
+ mean=flat.mean().item(),
144
+ max_abs=flat.abs().max().item(),
145
+ std=flat.std().item(),
146
+ norm=flat.norm().item(),
147
+ numel=flat.numel()
148
+ )
149
+ return stats
150
+
151
+
152
+ class LLMAnalyzer:
153
+ def __init__(
154
+ self,
155
+ device: Optional[str] = None,
156
+ max_length: int = 512,
157
+ max_samples: int = 1000, # number of wikitext samples to eval
158
+ batch_size: int = 4,
159
+ dtype: torch.dtype = torch.float16,
160
+ wandb_project: Optional[str] = None,
161
+ ):
162
+ self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")
163
+ self.max_length = max_length
164
+ self.max_samples = max_samples
165
+ self.batch_size = batch_size
166
+ self.dtype = dtype if self.device == "cuda" else torch.float32
167
+ self.wandb_project = wandb_project or os.environ.get("WANDB_PROJECT", "llm-activation-analysis")
168
+ self._cache = {}
169
+
170
+ def load_dataset(self, split: str = "test"):
171
+ """Load Salesforce/wikitext dataset."""
172
+ print(f"[Dataset] Loading Salesforce/wikitext ({split}) ...")
173
+ ds = load_dataset("Salesforce/wikitext", "wikitext-2-raw-v1", split=split)
174
+ # Filter out empty lines
175
+ texts = [t for t in ds["text"] if len(t.strip()) > 50]
176
+ print(f"[Dataset] Loaded {len(texts)} non-empty samples")
177
+ return texts[:self.max_samples]
178
+
179
+ def load_model(self, model_name: str):
180
+ """Load model and tokenizer with caching."""
181
+ if model_name in self._cache:
182
+ return self._cache[model_name]
183
+
184
+ print(f"[Loading] {model_name} ...")
185
+
186
+ tokenizer = AutoTokenizer.from_pretrained(
187
+ model_name,
188
+ trust_remote_code=True,
189
+ padding_side="left"
190
+ )
191
+ if tokenizer.pad_token is None:
192
+ tokenizer.pad_token = tokenizer.eos_token
193
+
194
+ config = AutoConfig.from_pretrained(model_name, trust_remote_code=True)
195
+
196
+ model = AutoModelForCausalLM.from_pretrained(
197
+ model_name,
198
+ config=config,
199
+ torch_dtype=self.dtype,
200
+ device_map="auto" if self.device == "cuda" else None,
201
+ trust_remote_code=True,
202
+ )
203
+
204
+ if self.device == "cpu":
205
+ model = model.to(self.device)
206
+
207
+ model.eval()
208
+
209
+ num_params = sum(p.numel() for p in model.parameters())
210
+
211
+ self._cache[model_name] = (tokenizer, model, config, num_params)
212
+ print(f"[Loaded] {model_name} | Params: {num_params/1e6:.1f}M | Layers: {config.num_hidden_layers} | Hidden: {config.hidden_size}")
213
+ return tokenizer, model, config, num_params
214
+
215
+ def compute(
216
+ self,
217
+ model_names: List[str],
218
+ ) -> List[ModelResult]:
219
+ """
220
+ Compute loss and per-tensor activation statistics for multiple models.
221
+ Logs each model as a separate WandB run.
222
+ """
223
+ texts = self.load_dataset()
224
+ results = []
225
+
226
+ for model_name in model_names:
227
+ try:
228
+ result = self._evaluate_model(model_name, texts)
229
+ results.append(result)
230
+ except Exception as e:
231
+ print(f"[Error] {model_name}: {e}")
232
+ import traceback
233
+ traceback.print_exc()
234
+ continue
235
+
236
+ return results
237
+
238
+ def _evaluate_model(
239
+ self,
240
+ model_name: str,
241
+ texts: List[str],
242
+ ) -> ModelResult:
243
+ tokenizer, model, config, num_params = self.load_model(model_name)
244
+
245
+ # Initialize WandB run for this model
246
+ run_name = model_name.replace("/", "-")
247
+ wandb.init(
248
+ project=self.wandb_project,
249
+ name=run_name,
250
+ config={
251
+ "model": model_name,
252
+ "max_length": self.max_length,
253
+ "max_samples": self.max_samples,
254
+ "batch_size": self.batch_size,
255
+ "dtype": str(self.dtype),
256
+ "num_params": num_params,
257
+ "num_layers": config.num_hidden_layers,
258
+ "hidden_size": config.hidden_size,
259
+ },
260
+ reinit=True
261
+ )
262
+
263
+ hook_mgr = ActivationHookManager()
264
+ hook_mgr.register_hooks(model)
265
+
266
+ total_loss = 0.0
267
+ total_tokens = 0
268
+
269
+ # Global activation accumulator
270
+ global_acts = []
271
+
272
+ # Per-layer activation accumulators
273
+ # We'll aggregate stats across batches, then compute final stats
274
+ layer_act_values = defaultdict(list)
275
+
276
+ num_batches = (len(texts) + self.batch_size - 1) // self.batch_size
277
+
278
+ for i in tqdm(range(0, len(texts), self.batch_size), desc=f"Eval {run_name}", total=num_batches):
279
+ batch_texts = texts[i:i + self.batch_size]
280
+
281
+ # Tokenize
282
+ inputs = tokenizer(
283
+ batch_texts,
284
+ return_tensors="pt",
285
+ truncation=True,
286
+ max_length=self.max_length,
287
+ padding=True
288
+ )
289
+
290
+ # Move to device
291
+ if self.device == "cuda" and hasattr(model, "device"):
292
+ # model is on auto device map
293
+ input_ids = inputs["input_ids"]
294
+ if hasattr(model, "device") and model.device != torch.device("meta"):
295
+ input_ids = input_ids.to(model.device)
296
+ attention_mask = inputs.get("attention_mask")
297
+ if attention_mask is not None:
298
+ attention_mask = attention_mask.to(input_ids.device)
299
+ else:
300
+ input_ids = inputs["input_ids"].to(self.device)
301
+ attention_mask = inputs.get("attention_mask")
302
+ if attention_mask is not None:
303
+ attention_mask = attention_mask.to(self.device)
304
+
305
+ labels = input_ids.clone()
306
+
307
+ # BUG FIX: with padding_side="left" and no explicit position_ids,
308
+ # HF's default `position_ids = arange(seq_len)` is applied
309
+ # identically to every row in the batch, regardless of how much
310
+ # left-padding precedes the real tokens in that row (verified
311
+ # against transformers' LlamaModel.forward / GPT2Model.forward
312
+ # source — neither adjusts for padding when position_ids=None).
313
+ # That means a real token's absolute position (and therefore its
314
+ # RoPE rotation / absolute position embedding) depends on how
315
+ # much padding happened to precede it in this particular batch,
316
+ # not on its logical position within its own sequence. This
317
+ # silently corrupts logits -> loss -> perplexity, with the
318
+ # amount of corruption varying batch-to-batch. Fix: derive
319
+ # position_ids from attention_mask so they restart at 0 for the
320
+ # first real token of every row, and are stable (0) on padding.
321
+ if attention_mask is not None:
322
+ position_ids = attention_mask.long().cumsum(-1) - 1
323
+ position_ids.masked_fill_(attention_mask == 0, 0)
324
+ else:
325
+ position_ids = None
326
+
327
+ hook_mgr.set_attention_mask(attention_mask)
328
+
329
+ with torch.no_grad():
330
+ outputs = model(
331
+ input_ids=input_ids,
332
+ attention_mask=attention_mask,
333
+ position_ids=position_ids,
334
+ labels=labels,
335
+ )
336
+
337
+ # --- Loss Computation ---
338
+ logits = outputs.logits
339
+ shift_logits = logits[..., :-1, :].contiguous()
340
+ shift_labels = labels[..., 1:].contiguous()
341
+ shift_mask = attention_mask[..., 1:].contiguous() if attention_mask is not None else None
342
+
343
+ loss_fct = torch.nn.CrossEntropyLoss(reduction="none")
344
+ token_losses = loss_fct(
345
+ shift_logits.view(-1, shift_logits.size(-1)),
346
+ shift_labels.view(-1)
347
+ )
348
+
349
+ if shift_mask is not None:
350
+ token_losses = token_losses * shift_mask.view(-1)
351
+ num_valid_tokens = shift_mask.sum().item()
352
+ else:
353
+ num_valid_tokens = token_losses.numel()
354
+
355
+ batch_loss = token_losses.sum().item()
356
+ total_loss += batch_loss
357
+ total_tokens += num_valid_tokens
358
+
359
+ # --- Activation Statistics ---
360
+ # Get activations captured by hooks
361
+ act_stats = hook_mgr.compute_stats()
362
+
363
+ for name, stats in act_stats.items():
364
+ # Skip dtype-limit fraction metrics entirely
365
+ # We only log mean, max_abs, std, norm
366
+
367
+ # For global: aggregate raw values
368
+ # We can't store all raw values due to memory, so we store running sums
369
+ # But for accurate std across all batches, we need a streaming algorithm
370
+ # For simplicity and correctness, we'll store per-batch stats and weight them
371
+ layer_act_values[name].append(asdict(stats))
372
+
373
+ hook_mgr.clear()
374
+
375
+ # Log per-batch metrics to wandb
376
+ if num_valid_tokens > 0:
377
+ batch_avg_loss = batch_loss / num_valid_tokens
378
+ wandb.log({
379
+ "batch_loss": batch_avg_loss,
380
+ "batch_perplexity": torch.exp(torch.tensor(batch_avg_loss)).item(),
381
+ "batch_tokens": num_valid_tokens,
382
+ "progress": i / len(texts)
383
+ }, step=i)
384
+
385
+ hook_mgr.remove_hooks()
386
+
387
+ # --- Final Aggregation ---
388
+ avg_loss = total_loss / max(total_tokens, 1)
389
+ perplexity = torch.exp(torch.tensor(avg_loss)).item()
390
+
391
+ # Aggregate per-tensor stats across all batches
392
+ # Weighted by numel for mean, max for max_abs, pooled std, pooled norm
393
+ final_layer_stats = {}
394
+
395
+ for name, batch_stats_list in layer_act_values.items():
396
+ total_numel = sum(s["numel"] for s in batch_stats_list)
397
+ if total_numel == 0:
398
+ continue
399
+
400
+ # Weighted mean
401
+ weighted_mean = sum(s["mean"] * s["numel"] for s in batch_stats_list) / total_numel
402
+
403
+ # Max abs across all batches
404
+ max_abs = max(s["max_abs"] for s in batch_stats_list)
405
+
406
+ # BUG FIX: the previous formula (weighted average of per-batch
407
+ # variances only) drops the between-batch term that accounts for
408
+ # per-batch means differing from the global mean. Whenever batch
409
+ # means differ (they will — different texts, different lengths),
410
+ # this systematically UNDERESTIMATES the true global std — in a
411
+ # quick numeric test with two batches of different means this was
412
+ # off by ~2.8x. Correct pooled-variance formula (population form,
413
+ # matches exp.py's StepAccumulator._merge_entry):
414
+ # E[X^2] = weighted_avg(var_i + mean_i^2)
415
+ # Var(X) = E[X^2] - mean_global^2
416
+ ex2 = sum(
417
+ s["numel"] * (s["std"] ** 2 + s["mean"] ** 2) for s in batch_stats_list
418
+ ) / total_numel
419
+ pooled_std = math.sqrt(max(0.0, ex2 - weighted_mean ** 2))
420
+
421
+ # Norm: sqrt(sum of squared norms / total_numel) * sqrt(total_numel)
422
+ # Actually norm^2 = sum(x_i^2), so pooled_norm = sqrt(sum(norm_i^2))
423
+ pooled_norm = (sum(s["norm"] ** 2 for s in batch_stats_list)) ** 0.5
424
+
425
+ final_layer_stats[name] = TensorStats(
426
+ mean=weighted_mean,
427
+ max_abs=max_abs,
428
+ std=pooled_std,
429
+ norm=pooled_norm,
430
+ numel=total_numel
431
+ )
432
+
433
+ # Compute global stats across all layers
434
+ if final_layer_stats:
435
+ all_numel = sum(s.numel for s in final_layer_stats.values())
436
+ global_mean = sum(s.mean * s.numel for s in final_layer_stats.values()) / all_numel
437
+ global_max_abs = max(s.max_abs for s in final_layer_stats.values())
438
+ # Same pooled-std correction as above — same bug was present here.
439
+ global_ex2 = sum(
440
+ s.numel * (s.std ** 2 + s.mean ** 2) for s in final_layer_stats.values()
441
+ ) / all_numel
442
+ global_std = math.sqrt(max(0.0, global_ex2 - global_mean ** 2))
443
+ global_norm = (sum(s.norm ** 2 for s in final_layer_stats.values())) ** 0.5
444
+
445
+ global_stats = TensorStats(
446
+ mean=global_mean,
447
+ max_abs=global_max_abs,
448
+ std=global_std,
449
+ norm=global_norm,
450
+ numel=all_numel
451
+ )
452
+ else:
453
+ global_stats = TensorStats(0.0, 0.0, 0.0, 0.0, 0)
454
+
455
+ result = ModelResult(
456
+ model_name=model_name,
457
+ loss=avg_loss,
458
+ perplexity=perplexity,
459
+ global_act=global_stats,
460
+ layer_acts=final_layer_stats,
461
+ num_tokens=total_tokens,
462
+ num_layers=config.num_hidden_layers,
463
+ hidden_size=config.hidden_size,
464
+ num_params=num_params
465
+ )
466
+
467
+ # --- WandB Logging ---
468
+ self._log_to_wandb(result)
469
+ wandb.finish()
470
+
471
+ return result
472
+
473
+ def _log_to_wandb(self, result: ModelResult):
474
+ """Log final metrics to WandB. No frac_near_dtype_limit."""
475
+
476
+ # Global metrics
477
+ wandb.log({
478
+ "final/loss": result.loss,
479
+ "final/perplexity": result.perplexity,
480
+ "final/num_tokens": result.num_tokens,
481
+
482
+ "train/global/act/mean": result.global_act.mean,
483
+ "train/global/act/max_abs": result.global_act.max_abs,
484
+ "train/global/act/std": result.global_act.std,
485
+ "train/global/act/norm": result.global_act.norm,
486
+ # Intentionally NOT logging frac_near_dtype_limit or frac_near_user_limit
487
+ })
488
+
489
+ # Per-tensor (per-layer) metrics
490
+ # Organize by layer for cleaner WandB UI
491
+ for tensor_name, stats in result.layer_acts.items():
492
+ # Clean name for wandb: replace dots with slashes
493
+ clean_name = tensor_name.replace(".", "/")
494
+
495
+ wandb.log({
496
+ f"train/{clean_name}/act/mean": stats.mean,
497
+ f"train/{clean_name}/act/max_abs": stats.max_abs,
498
+ f"train/{clean_name}/act/std": stats.std,
499
+ f"train/{clean_name}/act/norm": stats.norm,
500
+ # No frac_near_dtype_limit
501
+ })
502
+
503
+ # Also log as a wandb.Table for easy comparison
504
+ table_data = []
505
+ for tensor_name, stats in sorted(result.layer_acts.items()):
506
+ table_data.append([
507
+ tensor_name,
508
+ stats.mean,
509
+ stats.max_abs,
510
+ stats.std,
511
+ stats.norm,
512
+ stats.numel
513
+ ])
514
+
515
+ if table_data:
516
+ table = wandb.Table(
517
+ columns=["tensor_name", "mean", "max_abs", "std", "norm", "numel"],
518
+ data=table_data
519
+ )
520
+ wandb.log({"activation_table": table})
521
+
522
+ def print_report(self, results: List[ModelResult]):
523
+ """Pretty-print comparison report."""
524
+ print("\n" + "=" * 110)
525
+ print(f"{'Model':<35} {'Loss':>10} {'PPL':>10} {'ActMean':>12} {'ActMaxAbs':>12} {'ActStd':>12} {'Tokens':>8}")
526
+ print("-" * 110)
527
+
528
+ for r in results:
529
+ name = r.model_name.split("/")[-1][:33]
530
+ print(
531
+ f"{name:<35} "
532
+ f"{r.loss:>10.4f} "
533
+ f"{r.perplexity:>10.2f} "
534
+ f"{r.global_act.mean:>12.6f} "
535
+ f"{r.global_act.max_abs:>12.6f} "
536
+ f"{r.global_act.std:>12.6f} "
537
+ f"{r.num_tokens:>8}"
538
+ )
539
+
540
+ print("=" * 110)
541
+
542
+ # Print top 5 layers by max_abs for each model
543
+ print("\n[Per-Tensor Max Abs Top 5]")
544
+ for r in results:
545
+ name = r.model_name.split("/")[-1]
546
+ sorted_layers = sorted(r.layer_acts.items(), key=lambda x: x[1].max_abs, reverse=True)[:5]
547
+ print(f"\n {name}:")
548
+ for tensor_name, stats in sorted_layers:
549
+ print(f" {tensor_name:<50} max_abs={stats.max_abs:>10.4f} mean={stats.mean:>10.6f} std={stats.std:>10.4f}")
550
+
551
+ def export_json(self, results: List[ModelResult], path: str):
552
+ """Export results to JSON."""
553
+ data = []
554
+ for r in results:
555
+ entry = {
556
+ "model": r.model_name,
557
+ "loss": r.loss,
558
+ "perplexity": r.perplexity,
559
+ "num_tokens": r.num_tokens,
560
+ "num_layers": r.num_layers,
561
+ "hidden_size": r.hidden_size,
562
+ "num_params": r.num_params,
563
+ "global_act": asdict(r.global_act),
564
+ "layer_acts": {k: asdict(v) for k, v in r.layer_acts.items()}
565
+ }
566
+ data.append(entry)
567
+
568
+ with open(path, "w") as f:
569
+ json.dump(data, f, indent=2)
570
+ print(f"[Exported] Results saved to {path}")
zain/Activation/sweep.py ADDED
@@ -0,0 +1,217 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Sweep explicit GLU / MLP variants with identical data and hyperparameters."""
3
+
4
+ import argparse
5
+ import copy
6
+ import json
7
+ import re
8
+ import time
9
+ import os
10
+ from pathlib import Path
11
+
12
+ import yaml
13
+ import wandb
14
+ import torch
15
+ from transformers import AutoTokenizer, set_seed
16
+ from exp import (
17
+ TinyLlamaConfig,
18
+ TinyLlamaForCausalLM,
19
+ build_dataset,
20
+ create_trainer,
21
+ fetch_latest_checkpoint_from_hub,
22
+ ResumeFreezeOverrideCallback,
23
+ )
24
+
25
+
26
+ def format_param_count(total_params: int) -> str:
27
+ if total_params >= 1e9:
28
+ return f"{total_params / 1e9:.1f}B"
29
+ else:
30
+ return f"{total_params / 1e6:.1f}M"
31
+
32
+
33
+ def parse_variant(variant: str):
34
+ parts = variant.split('-')
35
+ if len(parts) < 2:
36
+ raise ValueError(f"Invalid variant format: '{variant}'. Expected: <glu|mlp>-<activation>[-<layers>L]")
37
+ prefix = parts[0]
38
+ if prefix not in ('glu', 'mlp'):
39
+ raise ValueError(f"Invalid prefix: '{prefix}'. Must be 'glu' or 'mlp'.")
40
+ last = parts[-1]
41
+ if last.endswith('L') and last[:-1].isdigit():
42
+ layers = int(last[:-1])
43
+ activation = '-'.join(parts[1:-1])
44
+ else:
45
+ layers = None
46
+ activation = '-'.join(parts[1:])
47
+ if not activation:
48
+ raise ValueError(f"Missing activation name in variant: '{variant}'")
49
+ return prefix, activation, layers
50
+
51
+
52
+ def main():
53
+ parser = argparse.ArgumentParser()
54
+ parser.add_argument("--config", required=True, help="Base YAML config")
55
+ parser.add_argument("--variants", nargs="+", required=True, help="List of variants")
56
+ parser.add_argument("--push", action="store_true")
57
+ args = parser.parse_args()
58
+
59
+ with open(args.config) as f:
60
+ base = yaml.safe_load(f)
61
+
62
+ seed = base.get("training", {}).get("seed", 42)
63
+ set_seed(seed)
64
+
65
+ wandb_project = base.get("training", {}).get("wandb_project")
66
+ if wandb_project:
67
+ os.environ["WANDB_PROJECT"] = wandb_project
68
+ print(f"[WandB] Project locked to: {wandb_project}")
69
+
70
+ tok_name = base["model"].get("tokenizer_name", "meta-llama/Llama-2-7b-hf")
71
+ tokenizer = AutoTokenizer.from_pretrained(tok_name)
72
+ if tokenizer.pad_token is None:
73
+ tokenizer.pad_token = tokenizer.eos_token
74
+
75
+ msl = base["model"].get("max_position_embeddings", 512)
76
+ train_ds = build_dataset(tokenizer, max_seq_len=msl, split="train", max_samples=None)
77
+ eval_ds = build_dataset(tokenizer, max_seq_len=msl, split="validation", max_samples=None)
78
+
79
+ results = []
80
+
81
+ for variant in args.variants:
82
+ prefix, act, layers = parse_variant(variant)
83
+
84
+ if prefix == "mlp" and act in ("situglu", "waleed", "situglu_low", "waleedglu_low"):
85
+ raise ValueError(
86
+ f"Activation '{act}' requires GLU. Please use 'glu-{act}'."
87
+ )
88
+
89
+ cfg = copy.deepcopy(base)
90
+ cfg["model"]["mlp_type"] = prefix
91
+ cfg["model"]["activation"] = act
92
+ if layers is not None:
93
+ cfg["model"]["num_hidden_layers"] = layers
94
+
95
+ actual_layers = cfg["model"]["num_hidden_layers"]
96
+ variant_label = f"{prefix}-{act}-{actual_layers}L"
97
+
98
+ # Output directory: add _trash if override enabled
99
+ go = cfg.get("gradient_override", {})
100
+ if go.get("enabled", False):
101
+ base_out = Path(cfg["training"]["output_dir"]).parent
102
+ variant_name = variant_label + "_trash_run"
103
+ else:
104
+ base_out = Path(cfg["training"]["output_dir"]).parent
105
+ variant_name = variant_label + "_run"
106
+
107
+ out_dir = base_out / variant_name
108
+ cfg["training"]["output_dir"] = str(out_dir)
109
+
110
+ # Checkpoint fetching
111
+ checkpoint_path = None
112
+ resume_cfg = cfg.get("training", {})
113
+ if resume_cfg.get("resume_from_hub", False):
114
+ hub_cfg = cfg.get("hub", {})
115
+ repo_id = hub_cfg.get("repo_id")
116
+ subpath = hub_cfg.get("subpath", "")
117
+ if not repo_id:
118
+ raise ValueError("hub.repo_id must be set when resume_from_hub is true")
119
+ checkpoint_step = resume_cfg.get("checkpoint_step")
120
+ print(f"[Resume] Fetching {variant_label} from HF Hub: {repo_id}/{subpath}/{variant_label}_run")
121
+ checkpoint_path = fetch_latest_checkpoint_from_hub(
122
+ repo_id=repo_id,
123
+ subpath=subpath,
124
+ variant=variant_label,
125
+ checkpoint_step=checkpoint_step,
126
+ )
127
+ print(f"[Resume] Downloaded to: {checkpoint_path}")
128
+ elif resume_cfg.get("resume_from"):
129
+ print("[Warning] resume_from set to a specific path; all variants will try to use the same checkpoint.")
130
+ checkpoint_path = resume_cfg.get("resume_from")
131
+
132
+ set_seed(seed)
133
+
134
+ print(f"\n{'='*60}\n>>> Variant: {variant_label} | Out: {out_dir}\n{'='*60}")
135
+
136
+ config = TinyLlamaConfig(**cfg["model"])
137
+ model = TinyLlamaForCausalLM(config)
138
+ model = model.to(torch.bfloat16)
139
+
140
+ total_params = sum(p.numel() for p in model.parameters())
141
+ param_str = format_param_count(total_params)
142
+ timestamp = time.strftime("%Y%m%d-%H%M%S")
143
+ run_name = f"LM-{variant_label}-{param_str}-{timestamp}"
144
+ cfg["training"]["run_name"] = run_name
145
+
146
+ hub_id_base = cfg["training"].get("hub_model_id", "tiny-llama-lab")
147
+ cfg["training"]["hub_model_id"] = f"{hub_id_base}-{variant_label}"
148
+
149
+ os.environ.pop("WANDB_RUN_ID", None)
150
+
151
+ trainer = create_trainer(model, tokenizer, cfg, train_ds, eval_ds)
152
+
153
+ # ---- CRITICAL FIX: manually set trainer on the callback ----
154
+ freeze_mlp = cfg["training"].get("freeze_mlp", False)
155
+ override_enabled = go.get("enabled", False)
156
+ if freeze_mlp or override_enabled:
157
+ for cb in trainer.callback_handler.callbacks:
158
+ if isinstance(cb, ResumeFreezeOverrideCallback):
159
+ cb._trainer = trainer
160
+ break
161
+
162
+ # Also delete optimizer.pt from the checkpoint so the Trainer doesn't try to load it
163
+ if checkpoint_path is not None and os.path.isdir(checkpoint_path):
164
+ opt_path = os.path.join(checkpoint_path, "optimizer.pt")
165
+ if os.path.exists(opt_path):
166
+ os.remove(opt_path)
167
+ print("[Freeze] Removed optimizer.pt from checkpoint to avoid state mismatch.")
168
+
169
+ # ---- NOW TRAIN ----
170
+ try:
171
+ trainer.train(resume_from_checkpoint=checkpoint_path)
172
+ metrics = trainer.evaluate()
173
+ results.append({
174
+ "variant": variant_label,
175
+ "eval_loss": metrics.get("eval_loss"),
176
+ "out": str(out_dir),
177
+ "run_name": run_name,
178
+ "status": "success",
179
+ })
180
+ trainer.save_model(str(out_dir))
181
+ if args.push or cfg["training"].get("push_to_hub", False):
182
+ trainer.push_to_hub()
183
+ print(f">>> FINISHED {variant_label} successfully")
184
+
185
+ except Exception as e:
186
+ print(f"!!! VARIANT {variant_label} FAILED with: {e}")
187
+ import traceback
188
+ traceback.print_exc()
189
+ results.append({
190
+ "variant": variant_label,
191
+ "error": str(e),
192
+ "out": str(out_dir),
193
+ "status": "failed",
194
+ })
195
+ # Continue with next variant
196
+ continue
197
+
198
+ finally:
199
+ wandb.finish()
200
+ torch.cuda.empty_cache() # free memory before next variant
201
+
202
+ summary = Path(base["training"]["output_dir"]).parent / "sweep_summary.json"
203
+ summary.write_text(json.dumps(results, indent=2))
204
+ print("\nSweep complete:")
205
+ for r in results:
206
+ if r.get("status") == "success":
207
+ print(f" {r['variant']:20s} eval_loss={r['eval_loss']:.4f}")
208
+ else:
209
+ print(f" {r['variant']:20s} FAILED: {r.get('error', 'unknown')}")
210
+
211
+ # Return non-zero exit code if any variant failed
212
+ if any(r.get("status") == "failed" for r in results):
213
+ sys.exit(1)
214
+
215
+
216
+ if __name__ == "__main__":
217
+ main()
zain/Activation/train.py ADDED
@@ -0,0 +1,122 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Train one TinyLlama variant with optional resume from HF Hub."""
3
+
4
+ import argparse
5
+ import yaml
6
+ import os
7
+ import torch
8
+ from pathlib import Path
9
+
10
+ from transformers import AutoTokenizer, set_seed
11
+ from exp import (
12
+ TinyLlamaConfig,
13
+ TinyLlamaForCausalLM,
14
+ build_dataset,
15
+ create_trainer,
16
+ fetch_latest_checkpoint_from_hub,
17
+ ResumeFreezeOverrideCallback,
18
+ )
19
+
20
+
21
+ def main():
22
+ parser = argparse.ArgumentParser()
23
+ parser.add_argument("--config", required=True, help="Path to YAML config")
24
+ parser.add_argument("--variant", required=True, help="e.g., glu-silu-100L")
25
+ parser.add_argument("--resume_from", help="Local checkpoint path (overrides config)")
26
+ parser.add_argument("--push", action="store_true", help="Push final model to HF Hub")
27
+ args = parser.parse_args()
28
+
29
+ with open(args.config) as f:
30
+ cfg = yaml.safe_load(f)
31
+
32
+ seed = cfg.get("training", {}).get("seed", 42)
33
+ set_seed(seed)
34
+
35
+ wandb_project = cfg.get("training", {}).get("wandb_project")
36
+ if wandb_project:
37
+ os.environ["WANDB_PROJECT"] = wandb_project
38
+ print(f"[WandB] Project locked to: {wandb_project}")
39
+
40
+ # Determine checkpoint path
41
+ checkpoint_path = None
42
+ if args.resume_from:
43
+ checkpoint_path = args.resume_from
44
+ print(f"[Resume] Using CLI-provided local path: {checkpoint_path}")
45
+ else:
46
+ resume_config = cfg.get("training", {})
47
+ if resume_config.get("resume_from_hub", False):
48
+ hub_cfg = cfg.get("hub", {})
49
+ repo_id = hub_cfg.get("repo_id")
50
+ subpath = hub_cfg.get("subpath", "")
51
+ if not repo_id:
52
+ raise ValueError("hub.repo_id must be set when resume_from_hub is true")
53
+ checkpoint_step = resume_config.get("checkpoint_step")
54
+ print(f"[Resume] Fetching from HF Hub: {repo_id}/{subpath}/{args.variant}_run")
55
+ checkpoint_path = fetch_latest_checkpoint_from_hub(
56
+ repo_id=repo_id,
57
+ subpath=subpath,
58
+ variant=args.variant,
59
+ checkpoint_step=checkpoint_step,
60
+ )
61
+ print(f"[Resume] Downloaded to: {checkpoint_path}")
62
+ elif resume_config.get("resume_from"):
63
+ checkpoint_path = resume_config.get("resume_from")
64
+ print(f"[Resume] Using config-provided local path: {checkpoint_path}")
65
+
66
+ # Build model & tokenizer
67
+ model_cfg = cfg["model"]
68
+ train_cfg = cfg.get("training", {})
69
+
70
+ # Output dir with _trash if override enabled
71
+ go = cfg.get("gradient_override", {})
72
+ if go.get("enabled", False):
73
+ base_out = Path(train_cfg.get("output_dir", "./out"))
74
+ variant_name = args.variant + "_trash"
75
+ output_dir = base_out / variant_name
76
+ train_cfg["output_dir"] = str(output_dir)
77
+ print(f"[Override] Output directory set to: {output_dir}")
78
+
79
+ tok_name = model_cfg.pop("tokenizer_name", "meta-llama/Llama-2-7b-hf")
80
+ tokenizer = AutoTokenizer.from_pretrained(tok_name)
81
+ if tokenizer.pad_token is None:
82
+ tokenizer.pad_token = tokenizer.eos_token
83
+
84
+ tiny_config = TinyLlamaConfig(**model_cfg)
85
+ model = TinyLlamaForCausalLM(tiny_config)
86
+ model = model.to(torch.bfloat16)
87
+
88
+ n_params = sum(p.numel() for p in model.parameters()) / 1e6
89
+ print(f"Model: {n_params:.2f}M params | MLP type: {tiny_config.mlp_type} | Activation: {tiny_config.activation}")
90
+
91
+ msl = model_cfg.get("max_position_embeddings", 512)
92
+ train_ds = build_dataset(tokenizer, max_seq_len=msl, split="train", max_samples=None)
93
+ eval_ds = build_dataset(tokenizer, max_seq_len=msl, split="validation", max_samples=None)
94
+
95
+ trainer = create_trainer(model, tokenizer, cfg, train_ds, eval_ds)
96
+
97
+ # Manually set trainer on callback if freezing/override is enabled
98
+ freeze_mlp = train_cfg.get("freeze_mlp", False)
99
+ override_enabled = go.get("enabled", False)
100
+ if freeze_mlp or override_enabled:
101
+ for cb in trainer.callback_handler.callbacks:
102
+ if isinstance(cb, ResumeFreezeOverrideCallback):
103
+ cb._trainer = trainer
104
+ break
105
+ # Delete optimizer.pt to avoid state mismatch
106
+ if checkpoint_path is not None and os.path.isdir(checkpoint_path):
107
+ opt_path = os.path.join(checkpoint_path, "optimizer.pt")
108
+ if os.path.exists(opt_path):
109
+ os.remove(opt_path)
110
+ print("[Freeze] Removed optimizer.pt from checkpoint.")
111
+
112
+ trainer.train(resume_from_checkpoint=checkpoint_path)
113
+
114
+ out = train_cfg.get("output_dir", "./out")
115
+ trainer.save_model(out)
116
+ if args.push or train_cfg.get("push_to_hub", False):
117
+ trainer.push_to_hub()
118
+ print(f"Done. Artifacts in {out}")
119
+
120
+
121
+ if __name__ == "__main__":
122
+ main()
zain/Activation/wandb/debug-internal.log ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ {"time":"2026-08-14T23:11:19.883985323Z","level":"INFO","msg":"wandb-core"}
2
+ {"time":"2026-08-14T23:11:19.88411539Z","level":"INFO","msg":"stream: starting","core version":"0.28.1"}
3
+ {"time":"2026-08-14T23:11:20.144460721Z","level":"INFO","msg":"stream: created new stream","id":"cvkq49dn"}
4
+ {"time":"2026-08-14T23:11:20.14455917Z","level":"INFO","msg":"handler: started"}
5
+ {"time":"2026-08-14T23:11:20.144636487Z","level":"INFO","msg":"stream: started"}
6
+ {"time":"2026-08-14T23:11:20.144648599Z","level":"INFO","msg":"writer: started","stream_id":"cvkq49dn"}
7
+ {"time":"2026-08-14T23:11:20.144706894Z","level":"INFO","msg":"sender: started"}
8
+ {"time":"2026-08-14T23:11:20.584339046Z","level":"INFO","msg":"filestream: sending request","total_files":1,"console_offset":0,"console_lines":1}
9
+ {"time":"2026-08-14T23:11:20.683540169Z","level":"INFO","msg":"filestream: request sent","status":"200 OK"}
zain/Activation/wandb/debug.log ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_setup.py:_flush():81] Current SDK version is 0.28.1
2
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_setup.py:_flush():81] Configure stats pid to 2164011
3
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_setup.py:_flush():81] Loading settings from environment variables
4
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_init.py:setup_run_log_directory():729] Logging user logs to /mnt/data/zainulabideen/zain-exp/notebooks/zain/Activation/wandb/run-20260814_231119-cvkq49dn/logs/debug.log
5
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_init.py:setup_run_log_directory():730] Logging internal logs to /mnt/data/zainulabideen/zain-exp/notebooks/zain/Activation/wandb/run-20260814_231119-cvkq49dn/logs/debug-internal.log
6
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_init.py:init():772] calling init triggers
7
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_init.py:init():777] wandb.init called with sweep_config: {}
8
+ config: {'_wandb': {}}
9
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_init.py:init():820] starting backend
10
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_init.py:init():826] Connected to an existing wandb-core service via WANDB_SERVICE
11
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_init.py:init():835] sending inform_init request
12
+ 2026-08-14 23:11:20,145 INFO MainThread:2164011 [wandb_init.py:init():840] backend started and connected
13
+ 2026-08-14 23:11:20,148 INFO MainThread:2164011 [wandb_init.py:init():910] updated telemetry
14
+ 2026-08-14 23:11:20,155 INFO MainThread:2164011 [wandb_init.py:init():933] communicating run to backend with 90.0 second timeout
15
+ 2026-08-14 23:11:20,497 INFO MainThread:2164011 [wandb_init.py:init():978] starting run threads in backend
16
+ 2026-08-14 23:11:20,570 INFO MainThread:2164011 [wandb_run.py:_console_start():2621] atexit reg
17
+ 2026-08-14 23:11:20,570 INFO MainThread:2164011 [wandb_run.py:_redirect():2471] redirect: wrap_raw
18
+ 2026-08-14 23:11:20,570 INFO MainThread:2164011 [wandb_run.py:_redirect():2540] Wrapping output streams.
19
+ 2026-08-14 23:11:20,570 INFO MainThread:2164011 [wandb_run.py:_redirect():2563] Redirects installed.
20
+ 2026-08-14 23:11:20,573 INFO MainThread:2164011 [wandb_init.py:init():1016] run started, returning control to user process
21
+ 2026-08-14 23:11:20,574 INFO MainThread:2164011 [wandb_run.py:_config_callback():1346] config_cb None None {'transformers_version': '5.16.0.dev0', 'architectures': None, 'output_hidden_states': False, 'return_dict': True, 'dtype': None, 'chunk_size_feed_forward': 0, 'is_encoder_decoder': False, 'id2label': {0: 'LABEL_0', 1: 'LABEL_1'}, 'label2id': {'LABEL_0': 0, 'LABEL_1': 1}, 'problem_type': None, 'vocab_size': 4096, 'hidden_size': 128, 'intermediate_size': 256, 'num_hidden_layers': 100, 'num_attention_heads': 4, 'num_key_value_heads': 4, 'hidden_act': 'silu', 'max_position_embeddings': 512, 'initializer_range': 0.02, 'rms_norm_eps': 1e-06, 'use_cache': False, 'pad_token_id': 0, 'bos_token_id': 1, 'eos_token_id': 2, 'pretraining_tp': 1, 'tie_word_embeddings': True, 'rope_parameters': {'rope_theta': 10000.0, 'rope_type': 'default'}, 'attention_bias': False, 'attention_dropout': 0.0, 'mlp_bias': False, 'head_dim': 32, '_name_or_path': '', 'tokenizer_name': 'w-ahmad/tiny-stories-tokenizer', 'mlp_type': 'glu', 'activation': 'waleed10', 'waleed_beta': 2, 'model_type': 'tiny_llama', 'output_attentions': False, 'output_dir': 'out/glu-waleed10-100L_trash_run', 'per_device_train_batch_size': 64, 'num_train_epochs': 1, 'max_steps': 3500, 'learning_rate': 9e-05, 'lr_scheduler_type': 'constant_with_warmup', 'lr_scheduler_kwargs': None, 'warmup_steps': 50, 'optim': 'adamw_torch_fused', 'optim_args': None, 'weight_decay': 0.0, 'adam_beta1': 0.9, 'adam_beta2': 0.999, 'adam_epsilon': 1e-08, 'optim_target_modules': None, 'gradient_accumulation_steps': 1, 'average_tokens_across_devices': True, 'max_grad_norm': 1.0, 'label_smoothing_factor': 0.0, 'bf16': True, 'fp16': False, 'bf16_full_eval': False, 'fp16_full_eval': False, 'tf32': None, 'gradient_checkpointing': False, 'gradient_checkpointing_kwargs': None, 'torch_compile': False, 'torch_compile_backend': None, 'torch_compile_mode': None, 'use_liger_kernel': False, 'liger_kernel_config': None, 'neftune_noise_alpha': None, 'torch_empty_cache_steps': None, 'auto_find_batch_size': False, 'logging_strategy': 'steps', 'logging_steps': 20, 'logging_first_step': False, 'log_on_each_node': True, 'logging_nan_inf_filter': True, 'include_num_input_tokens_seen': 'no', 'log_level': 'passive', 'log_level_replica': 'warning', 'disable_tqdm': False, 'report_to': ['wandb'], 'run_name': 'LM-glu-waleed10-100L-16.9M-20260814-231118', 'project': 'huggingface', 'trackio_space_id': None, 'trackio_bucket_id': None, 'trackio_static_space_id': None, 'eval_strategy': 'steps', 'eval_steps': 2498, 'eval_delay': 0, 'per_device_eval_batch_size': 1024, 'prediction_loss_only': False, 'eval_on_start': False, 'eval_do_concat_batches': True, 'eval_use_gather_object': False, 'eval_accumulation_steps': None, 'include_for_metrics': [], 'batch_eval_metrics': False, 'save_only_model': False, 'save_strategy': 'steps', 'save_steps': 50, 'save_on_each_node': False, 'save_total_limit': None, 'enable_jit_checkpoint': False, 'push_to_hub': False, 'hub_token': '<HUB_TOKEN>', 'hub_private_repo': None, 'hub_model_id': 'w-ahmad/6L-glu-waleed10-100L', 'hub_strategy': 'every_save', 'hub_always_push': False, 'hub_revision': None, 'load_best_model_at_end': False, 'metric_for_best_model': None, 'greater_is_better': None, 'ignore_data_skip': False, 'restore_callback_states_from_checkpoint': False, 'full_determinism': False, 'seed': 42, 'data_seed': 42, 'use_cpu': False, 'accelerator_config': {'split_batches': False, 'dispatch_batches': None, 'even_batches': True, 'use_seedable_sampler': True, 'non_blocking': False, 'gradient_accumulation_kwargs': None}, 'parallelism_config': None, 'dataloader_drop_last': False, 'dataloader_num_workers': 0, 'dataloader_pin_memory': True, 'dataloader_persistent_workers': False, 'dataloader_prefetch_factor': None, 'dataloader_multiprocessing_context': None, 'dataloader_in_order': True, 'remove_unused_columns': False, 'label_names': None, 'train_sampling_strategy': 'random', 'length_column_name': 'length', 'ddp_find_unused_parameters': None, 'ddp_bucket_cap_mb': None, 'ddp_broadcast_buffers': None, 'ddp_static_graph': None, 'ddp_backend': None, 'ddp_timeout': 1800, 'fsdp': None, 'fsdp_config': None, 'deepspeed': None, 'debug': [], 'skip_memory_metrics': True, 'do_train': False, 'do_eval': True, 'do_predict': False, 'resume_from_checkpoint': None, 'local_rank': -1}
22
+ 2026-08-14 23:11:20,578 INFO MainThread:2164011 [wandb_config.py:__setitem__():155] [no run ID] config set model/num_parameters = 16934016 - <bound method Run._config_callback of <wandb.sdk.wandb_run.Run object at 0x14c6ee183d10>>
23
+ 2026-08-14 23:11:20,578 INFO MainThread:2164011 [wandb_run.py:_config_callback():1346] config_cb model/num_parameters 16934016 None
zain/Activation/wandb/run-20260814_231119-cvkq49dn/files/output.log ADDED
@@ -0,0 +1 @@
 
 
1
+ [Override] Activated with value = -10.0
zain/Activation/wandb/run-20260814_231119-cvkq49dn/files/requirements.txt ADDED
@@ -0,0 +1,149 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ asttokens==3.0.1
2
+ comm==0.2.3
3
+ debugpy==1.8.21
4
+ decorator==5.3.1
5
+ executing==2.2.1
6
+ nest-asyncio==1.6.0
7
+ parso==0.8.7
8
+ platformdirs==4.11.0
9
+ psutil==7.2.2
10
+ ptyprocess==0.7.0
11
+ pure_eval==0.2.3
12
+ Pygments==2.20.0
13
+ pyzmq==27.1.0
14
+ setuptools==83.0.0
15
+ six==1.17.0
16
+ tornado==6.5.7
17
+ traitlets==5.15.0
18
+ fsspec==2026.4.0
19
+ wcwidth==0.8.2
20
+ ipython_pygments_lexers==1.1.1
21
+ jedi==0.20.0
22
+ jupyter_core==5.9.1
23
+ matplotlib-inline==0.2.2
24
+ pexpect==4.9.0
25
+ prompt_toolkit==3.0.53
26
+ python-dateutil==2.9.0.post0
27
+ stack_data==0.6.3
28
+ wheel==0.47.0
29
+ jupyter_client==8.9.1
30
+ pip==26.1.2
31
+ ipython==9.15.0
32
+ ipykernel==7.2.0
33
+ threadpoolctl==3.6.0
34
+ pyparsing==3.3.2
35
+ typing_extensions==4.15.0
36
+ Jinja2==3.1.6
37
+ narwhals==2.24.0
38
+ kiwisolver==1.5.0
39
+ joblib==1.5.3
40
+ fonttools==4.63.0
41
+ cycler==0.12.1
42
+ scipy==1.17.1
43
+ pandas==3.0.5
44
+ contourpy==1.3.3
45
+ scikit-learn==1.9.0
46
+ matplotlib==3.11.1
47
+ urllib3==2.7.0
48
+ tqdm==4.70.0
49
+ idna==3.18
50
+ charset-normalizer==3.4.9
51
+ certifi==2026.7.22
52
+ requests==2.34.2
53
+ seaborn==0.13.2
54
+ uv==0.12.0
55
+ shellingham==1.5.4
56
+ mpmath==1.3.0
57
+ attrs==26.1.0
58
+ hf-xet==1.5.2
59
+ nvidia-nccl-cu12==2.21.5
60
+ MarkupSafe==3.0.3
61
+ regex==2026.7.19
62
+ importlib_metadata==9.0.0
63
+ httpcore==1.0.9
64
+ annotated-doc==0.0.5
65
+ multidict==6.7.1
66
+ aiohttp==3.14.3
67
+ aiosignal==1.4.0
68
+ xxhash==3.8.1
69
+ aiohappyeyeballs==2.7.1
70
+ mdurl==0.1.2
71
+ cuda-toolkit==13.0.3.0
72
+ networkx==3.6.1
73
+ PyYAML==6.0.3
74
+ nvidia-cufile==1.15.1.6
75
+ typer==0.27.0
76
+ torchaudio==2.6.0+cu124
77
+ rich==15.0.0
78
+ nvidia-cufft-cu12==11.2.1.3
79
+ h11==0.16.0
80
+ dill==0.4.1
81
+ cuda-pathfinder==1.6.0
82
+ filelock==3.29.0
83
+ nvidia-nvtx-cu12==12.4.127
84
+ httpx==0.28.1
85
+ anyio==4.14.2
86
+ numpy==2.4.4
87
+ yarl==1.24.5
88
+ click==8.4.2
89
+ triton==3.2.0
90
+ frozenlist==1.8.0
91
+ zipp==4.1.0
92
+ propcache==0.5.2
93
+ tokenizers==0.22.2
94
+ markdown-it-py==4.2.0
95
+ nvidia-cuda-runtime==13.0.96
96
+ cuda-bindings==13.3.1
97
+ nvidia-cuda-cupti==13.0.85
98
+ torch==2.6.0+cu124
99
+ multiprocess==0.70.19
100
+ pillow==12.2.0
101
+ transformers==5.16.0.dev0
102
+ wandb==0.28.1
103
+ nvidia-curand==10.4.0.35
104
+ sympy==1.13.1
105
+ nvidia-cusparse==12.6.3.3
106
+ nvidia-cuda-nvrtc==13.0.88
107
+ typing-inspection==0.4.2
108
+ nvidia-cusolver==12.0.4.66
109
+ nvidia-cufft==12.0.0.61
110
+ nvidia-cudnn-cu13==9.20.0.48
111
+ nvidia-cublas==13.1.1.3
112
+ pyarrow==25.0.0
113
+ evaluate==0.4.6
114
+ diffusers==0.39.0
115
+ pydantic==2.13.4
116
+ annotated-types==0.8.0
117
+ protobuf==7.35.1
118
+ sentry-sdk==2.66.1
119
+ einops==0.8.2
120
+ packaging==26.2
121
+ nvidia-nvjitlink-cu12==12.4.127
122
+ nvidia-curand-cu12==10.3.5.147
123
+ nvidia-cusparselt-cu12==0.6.2
124
+ nvidia-cusparse-cu12==12.3.1.170
125
+ nvidia-cuda-runtime-cu12==12.4.127
126
+ torchvision==0.21.0+cu124
127
+ nvidia-cuda-nvrtc-cu12==12.4.127
128
+ nvidia-cuda-cupti-cu12==12.4.127
129
+ nvidia-cusolver-cu12==11.6.1.9
130
+ nvidia-cublas-cu12==12.4.5.8
131
+ nvidia-cudnn-cu12==9.1.0.70
132
+ huggingface_hub==1.26.0
133
+ datasets==5.0.1
134
+ safetensors==0.8.0
135
+ accelerate==1.14.0
136
+ pydantic_core==2.46.4
137
+ ninja==1.13.0
138
+ autocommand==2.2.2
139
+ backports.tarfile==1.2.0
140
+ importlib_metadata==8.7.1
141
+ jaraco.text==4.0.0
142
+ jaraco.context==6.1.0
143
+ jaraco.functools==4.4.0
144
+ more-itertools==10.8.0
145
+ packaging==26.0
146
+ platformdirs==4.4.0
147
+ tomli==2.4.0
148
+ wheel==0.46.3
149
+ zipp==3.23.0
zain/Activation/wandb/run-20260814_231119-cvkq49dn/files/wandb-metadata.json ADDED
@@ -0,0 +1,96 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "os": "Linux-5.15.0-126-generic-x86_64-with-glibc2.35",
3
+ "python": "CPython 3.11.15",
4
+ "startedAt": "2026-08-14T23:11:19.881495Z",
5
+ "args": [
6
+ "--config",
7
+ "/mnt/data/zainulabideen/zain-exp/notebooks/zain/Activation/configs/baseline100l.yaml",
8
+ "--variants",
9
+ "glu-waleed10",
10
+ "glu-silu-waleed10"
11
+ ],
12
+ "program": "/mnt/data/zainulabideen/zain-exp/notebooks/zain/Activation/sweep.py",
13
+ "codePath": "sweep.py",
14
+ "codePathLocal": "sweep.py",
15
+ "git": {
16
+ "remote": "https://github.com/w-ahmad1a10/Activation.git",
17
+ "commit": "3fb47c6943d6927eafc39bd399a3d3a520bae74a"
18
+ },
19
+ "email": "deepnevro@gmail.com",
20
+ "root": "/mnt/data/zainulabideen/zain-exp/notebooks/zain/Activation",
21
+ "host": "deeplens-k3s-node1",
22
+ "executable": "/mnt/data/zainulabideen/zain-exp/notebooks/my_env/bin/python",
23
+ "cpu_count": 112,
24
+ "cpu_count_logical": 224,
25
+ "gpu": "NVIDIA H100 80GB HBM3",
26
+ "gpu_count": 8,
27
+ "disk": {
28
+ "/": {
29
+ "total": "1560765693952",
30
+ "used": "709375602688"
31
+ }
32
+ },
33
+ "memory": {
34
+ "total": "2164089937920"
35
+ },
36
+ "gpu_nvidia": [
37
+ {
38
+ "name": "NVIDIA H100 80GB HBM3",
39
+ "memoryTotal": "85520809984",
40
+ "cudaCores": 16896,
41
+ "architecture": "Hopper",
42
+ "uuid": "GPU-39c684a5-fde6-83d7-1663-0859795881ae"
43
+ },
44
+ {
45
+ "name": "NVIDIA H100 80GB HBM3",
46
+ "memoryTotal": "85520809984",
47
+ "cudaCores": 16896,
48
+ "architecture": "Hopper",
49
+ "uuid": "GPU-68012e5a-38b6-b643-0ca6-62fb66720bf3"
50
+ },
51
+ {
52
+ "name": "NVIDIA H100 80GB HBM3",
53
+ "memoryTotal": "85520809984",
54
+ "cudaCores": 16896,
55
+ "architecture": "Hopper",
56
+ "uuid": "GPU-132944c4-b689-2b5f-89a4-d730401677ab"
57
+ },
58
+ {
59
+ "name": "NVIDIA H100 80GB HBM3",
60
+ "memoryTotal": "85520809984",
61
+ "cudaCores": 16896,
62
+ "architecture": "Hopper",
63
+ "uuid": "GPU-2df386cc-6d26-d0e2-7a2d-a057b0d95864"
64
+ },
65
+ {
66
+ "name": "NVIDIA H100 80GB HBM3",
67
+ "memoryTotal": "85520809984",
68
+ "cudaCores": 16896,
69
+ "architecture": "Hopper",
70
+ "uuid": "GPU-bfa16575-1d94-1aa2-4537-2c93433f42ef"
71
+ },
72
+ {
73
+ "name": "NVIDIA H100 80GB HBM3",
74
+ "memoryTotal": "85520809984",
75
+ "cudaCores": 16896,
76
+ "architecture": "Hopper",
77
+ "uuid": "GPU-bc6c3e3c-9b90-09ca-c034-774961847c54"
78
+ },
79
+ {
80
+ "name": "NVIDIA H100 80GB HBM3",
81
+ "memoryTotal": "85520809984",
82
+ "cudaCores": 16896,
83
+ "architecture": "Hopper",
84
+ "uuid": "GPU-00a441e1-7c95-e7d6-4c35-43d6b291aea9"
85
+ },
86
+ {
87
+ "name": "NVIDIA H100 80GB HBM3",
88
+ "memoryTotal": "85520809984",
89
+ "cudaCores": 16896,
90
+ "architecture": "Hopper",
91
+ "uuid": "GPU-1c4d29a2-4647-6fce-d8fc-0c5ecfbbd6ea"
92
+ }
93
+ ],
94
+ "cudaVersion": "12.4",
95
+ "writerId": "f28mx2dyagzawnyc39dp62mh4jx7ni7b"
96
+ }
zain/Activation/wandb/run-20260814_231119-cvkq49dn/logs/debug-core.log ADDED
@@ -0,0 +1,88 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {"time":"2026-08-14T19:38:51.636058907Z","level":"INFO","msg":"main: starting server","port-filename":"/tmp/tmp54729boe/port-1375539.txt","pid":1375539,"detached":false,"idle-timeout":600000000000,"log-level":0,"disable-analytics":false,"shutdown-on-parent-exit":false,"enable-dcgm-profiling":false}
2
+ {"time":"2026-08-14T19:38:51.639970993Z","level":"INFO","msg":"server: will exit if parent process dies","ppid":1375539}
3
+ {"time":"2026-08-14T19:38:51.639940209Z","level":"INFO","msg":"server: accepting connections","addr":{"Name":"/tmp/wandb-1375539-170701-251437063/socket","Net":"unix"}}
4
+ {"time":"2026-08-14T19:38:51.812900736Z","level":"INFO","msg":"connection: ManageConnectionData: new connection created","id":"1(@)"}
5
+ {"time":"2026-08-14T19:39:07.276688434Z","level":"INFO","msg":"connection: ManageConnectionData: new connection created","id":"2(@)"}
6
+ {"time":"2026-08-14T19:39:07.370077387Z","level":"INFO","msg":"handleInformInit: received","streamId":"kk1xbqht","id":"2(@)"}
7
+ {"time":"2026-08-14T19:39:07.640206016Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"kk1xbqht","id":"2(@)"}
8
+ {"time":"2026-08-14T19:39:13.074249733Z","level":"INFO","msg":"connection: cancelling request","id":"2(@)","requestId":"zj3lg1k64oen"}
9
+ {"time":"2026-08-14T19:51:49.732364673Z","level":"INFO","msg":"connection: cancelling request","id":"2(@)","requestId":"zj3lg1k64oen"}
10
+ {"time":"2026-08-14T19:51:50.223490459Z","level":"INFO","msg":"connection: cancelling request","id":"2(@)","requestId":"zj3lg1k64oen"}
11
+ {"time":"2026-08-14T19:51:50.225275352Z","level":"INFO","msg":"handleInformFinish: finish message received","streamId":"kk1xbqht","id":"2(@)"}
12
+ {"time":"2026-08-14T19:51:50.225791347Z","level":"INFO","msg":"handleInformFinish: stream closed","streamId":"kk1xbqht","id":"2(@)"}
13
+ {"time":"2026-08-14T19:51:51.140595835Z","level":"INFO","msg":"handleInformInit: received","streamId":"ix6wcbwv","id":"2(@)"}
14
+ {"time":"2026-08-14T19:51:51.407779103Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"ix6wcbwv","id":"2(@)"}
15
+ {"time":"2026-08-14T19:51:56.793135206Z","level":"INFO","msg":"connection: cancelling request","id":"2(@)","requestId":"m2u1nqrk4r4j"}
16
+ {"time":"2026-08-14T20:05:01.389026301Z","level":"INFO","msg":"connection: cancelling request","id":"2(@)","requestId":"m2u1nqrk4r4j"}
17
+ {"time":"2026-08-14T20:05:01.891785946Z","level":"INFO","msg":"connection: cancelling request","id":"2(@)","requestId":"m2u1nqrk4r4j"}
18
+ {"time":"2026-08-14T20:05:01.893356932Z","level":"INFO","msg":"handleInformFinish: finish message received","streamId":"ix6wcbwv","id":"2(@)"}
19
+ {"time":"2026-08-14T20:05:01.894067235Z","level":"INFO","msg":"handleInformFinish: stream closed","streamId":"ix6wcbwv","id":"2(@)"}
20
+ {"time":"2026-08-14T20:05:03.038763184Z","level":"INFO","msg":"handleInformInit: received","streamId":"bdhnr27j","id":"2(@)"}
21
+ {"time":"2026-08-14T20:05:03.299131712Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"bdhnr27j","id":"2(@)"}
22
+ {"time":"2026-08-14T20:05:08.644597551Z","level":"INFO","msg":"connection: cancelling request","id":"2(@)","requestId":"u5vxhmn1tauv"}
23
+ {"time":"2026-08-14T20:21:26.598738734Z","level":"INFO","msg":"connection: cancelling request","id":"2(@)","requestId":"u5vxhmn1tauv"}
24
+ {"time":"2026-08-14T20:21:27.243676424Z","level":"INFO","msg":"connection: cancelling request","id":"2(@)","requestId":"u5vxhmn1tauv"}
25
+ {"time":"2026-08-14T20:21:27.245262581Z","level":"INFO","msg":"handleInformFinish: finish message received","streamId":"bdhnr27j","id":"2(@)"}
26
+ {"time":"2026-08-14T20:21:27.245859065Z","level":"INFO","msg":"handleInformFinish: stream closed","streamId":"bdhnr27j","id":"2(@)"}
27
+ {"time":"2026-08-14T20:21:28.155191217Z","level":"INFO","msg":"handleInformInit: received","streamId":"twgo1e89","id":"2(@)"}
28
+ {"time":"2026-08-14T20:21:28.414618766Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"twgo1e89","id":"2(@)"}
29
+ {"time":"2026-08-14T20:21:33.76572887Z","level":"INFO","msg":"connection: cancelling request","id":"2(@)","requestId":"smlkpmpk76rp"}
30
+ {"time":"2026-08-14T20:38:01.401128169Z","level":"INFO","msg":"connection: cancelling request","id":"2(@)","requestId":"smlkpmpk76rp"}
31
+ {"time":"2026-08-14T20:38:02.166866415Z","level":"INFO","msg":"connection: cancelling request","id":"2(@)","requestId":"smlkpmpk76rp"}
32
+ {"time":"2026-08-14T20:38:02.168373578Z","level":"INFO","msg":"handleInformFinish: finish message received","streamId":"twgo1e89","id":"2(@)"}
33
+ {"time":"2026-08-14T20:38:02.169204491Z","level":"INFO","msg":"handleInformFinish: stream closed","streamId":"twgo1e89","id":"2(@)"}
34
+ {"time":"2026-08-14T20:38:03.474631884Z","level":"INFO","msg":"handleInformInit: received","streamId":"7p5140bq","id":"2(@)"}
35
+ {"time":"2026-08-14T20:38:03.739404113Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"7p5140bq","id":"2(@)"}
36
+ {"time":"2026-08-14T20:38:09.1668666Z","level":"INFO","msg":"connection: cancelling request","id":"2(@)","requestId":"9gr96sam7ydd"}
37
+ {"time":"2026-08-14T20:54:40.077063626Z","level":"INFO","msg":"connection: cancelling request","id":"2(@)","requestId":"9gr96sam7ydd"}
38
+ {"time":"2026-08-14T20:54:40.77979146Z","level":"INFO","msg":"connection: cancelling request","id":"2(@)","requestId":"9gr96sam7ydd"}
39
+ {"time":"2026-08-14T20:54:40.781496498Z","level":"INFO","msg":"handleInformFinish: finish message received","streamId":"7p5140bq","id":"2(@)"}
40
+ {"time":"2026-08-14T20:54:40.782279361Z","level":"INFO","msg":"handleInformFinish: stream closed","streamId":"7p5140bq","id":"2(@)"}
41
+ {"time":"2026-08-14T20:54:42.400855251Z","level":"INFO","msg":"connection: closing","id":"2(@)"}
42
+ {"time":"2026-08-14T20:54:42.40094641Z","level":"INFO","msg":"connection: closed successfully","id":"2(@)"}
43
+ {"time":"2026-08-14T20:54:42.400864953Z","level":"INFO","msg":"processOutgoingData: finished","id":"2(@)"}
44
+ {"time":"2026-08-14T20:54:42.400960827Z","level":"INFO","msg":"connection: ManageConnectionData: connection closed","id":"2(@)"}
45
+ {"time":"2026-08-14T21:01:09.497626422Z","level":"INFO","msg":"connection: ManageConnectionData: new connection created","id":"3(@)"}
46
+ {"time":"2026-08-14T21:01:09.714752804Z","level":"INFO","msg":"handleInformInit: received","streamId":"qauftiuh","id":"3(@)"}
47
+ {"time":"2026-08-14T21:01:09.994191533Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"qauftiuh","id":"3(@)"}
48
+ {"time":"2026-08-14T21:01:15.507216386Z","level":"INFO","msg":"connection: cancelling request","id":"3(@)","requestId":"wd49dhcig506"}
49
+ {"time":"2026-08-14T21:07:53.011770412Z","level":"INFO","msg":"connection: cancelling request","id":"3(@)","requestId":"wd49dhcig506"}
50
+ {"time":"2026-08-14T21:07:54.661515353Z","level":"INFO","msg":"connection: cancelling request","id":"3(@)","requestId":"wd49dhcig506"}
51
+ {"time":"2026-08-14T21:07:54.926288764Z","level":"INFO","msg":"handleInformFinish: finish message received","streamId":"qauftiuh","id":"3(@)"}
52
+ {"time":"2026-08-14T21:07:54.927423677Z","level":"INFO","msg":"handleInformFinish: stream closed","streamId":"qauftiuh","id":"3(@)"}
53
+ {"time":"2026-08-14T21:08:09.037499672Z","level":"INFO","msg":"handleInformInit: received","streamId":"a6574zx8","id":"3(@)"}
54
+ {"time":"2026-08-14T21:08:09.296770198Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"a6574zx8","id":"3(@)"}
55
+ {"time":"2026-08-14T21:08:09.806890437Z","level":"INFO","msg":"connection: cancelling request","id":"3(@)","requestId":"bsqpd0c29zk9"}
56
+ {"time":"2026-08-14T21:08:14.888124439Z","level":"INFO","msg":"connection: cancelling request","id":"3(@)","requestId":"ilk9hnwtzxxz"}
57
+ {"time":"2026-08-14T21:14:51.274587185Z","level":"INFO","msg":"connection: cancelling request","id":"3(@)","requestId":"ilk9hnwtzxxz"}
58
+ {"time":"2026-08-14T21:14:52.85240857Z","level":"INFO","msg":"connection: cancelling request","id":"3(@)","requestId":"ilk9hnwtzxxz"}
59
+ {"time":"2026-08-14T21:14:52.873645397Z","level":"INFO","msg":"handleInformFinish: finish message received","streamId":"a6574zx8","id":"3(@)"}
60
+ {"time":"2026-08-14T21:14:52.874311484Z","level":"INFO","msg":"handleInformFinish: stream closed","streamId":"a6574zx8","id":"3(@)"}
61
+ {"time":"2026-08-14T21:15:07.169254211Z","level":"INFO","msg":"handleInformInit: received","streamId":"54wnyexz","id":"3(@)"}
62
+ {"time":"2026-08-14T21:15:07.435699541Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"54wnyexz","id":"3(@)"}
63
+ {"time":"2026-08-14T21:15:12.853675736Z","level":"INFO","msg":"connection: cancelling request","id":"3(@)","requestId":"fzqy9cymfofj"}
64
+ {"time":"2026-08-14T21:21:52.842687654Z","level":"INFO","msg":"connection: cancelling request","id":"3(@)","requestId":"fzqy9cymfofj"}
65
+ {"time":"2026-08-14T21:21:54.009681636Z","level":"INFO","msg":"connection: cancelling request","id":"3(@)","requestId":"fzqy9cymfofj"}
66
+ {"time":"2026-08-14T21:21:54.257319168Z","level":"INFO","msg":"handleInformFinish: finish message received","streamId":"54wnyexz","id":"3(@)"}
67
+ {"time":"2026-08-14T21:21:54.257934554Z","level":"INFO","msg":"handleInformFinish: stream closed","streamId":"54wnyexz","id":"3(@)"}
68
+ {"time":"2026-08-14T21:22:08.718181031Z","level":"INFO","msg":"handleInformInit: received","streamId":"n6fjpg9o","id":"3(@)"}
69
+ {"time":"2026-08-14T21:22:08.976795469Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"n6fjpg9o","id":"3(@)"}
70
+ {"time":"2026-08-14T21:22:14.472169125Z","level":"INFO","msg":"connection: cancelling request","id":"3(@)","requestId":"qlc6t75fdq0u"}
71
+ {"time":"2026-08-14T21:28:18.727248774Z","level":"INFO","msg":"connection: cancelling request","id":"3(@)","requestId":"qlc6t75fdq0u"}
72
+ {"time":"2026-08-14T21:28:20.378913364Z","level":"INFO","msg":"connection: cancelling request","id":"3(@)","requestId":"qlc6t75fdq0u"}
73
+ {"time":"2026-08-14T21:28:20.400345521Z","level":"INFO","msg":"handleInformFinish: finish message received","streamId":"n6fjpg9o","id":"3(@)"}
74
+ {"time":"2026-08-14T21:28:20.400969576Z","level":"INFO","msg":"handleInformFinish: stream closed","streamId":"n6fjpg9o","id":"3(@)"}
75
+ {"time":"2026-08-14T21:28:35.344777056Z","level":"INFO","msg":"handleInformInit: received","streamId":"gqskne0x","id":"3(@)"}
76
+ {"time":"2026-08-14T21:28:35.603024157Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"gqskne0x","id":"3(@)"}
77
+ {"time":"2026-08-14T21:28:41.078893683Z","level":"INFO","msg":"connection: cancelling request","id":"3(@)","requestId":"yxbo55z35lse"}
78
+ {"time":"2026-08-14T21:35:05.007052341Z","level":"INFO","msg":"connection: cancelling request","id":"3(@)","requestId":"yxbo55z35lse"}
79
+ {"time":"2026-08-14T21:35:06.745707472Z","level":"INFO","msg":"connection: cancelling request","id":"3(@)","requestId":"yxbo55z35lse"}
80
+ {"time":"2026-08-14T21:35:06.767299957Z","level":"INFO","msg":"handleInformFinish: finish message received","streamId":"gqskne0x","id":"3(@)"}
81
+ {"time":"2026-08-14T21:35:06.773662133Z","level":"INFO","msg":"handleInformFinish: stream closed","streamId":"gqskne0x","id":"3(@)"}
82
+ {"time":"2026-08-14T21:35:08.806158062Z","level":"INFO","msg":"connection: closing","id":"3(@)"}
83
+ {"time":"2026-08-14T21:35:08.806241813Z","level":"INFO","msg":"connection: closed successfully","id":"3(@)"}
84
+ {"time":"2026-08-14T21:35:08.806161779Z","level":"INFO","msg":"processOutgoingData: finished","id":"3(@)"}
85
+ {"time":"2026-08-14T21:35:08.806251477Z","level":"INFO","msg":"connection: ManageConnectionData: connection closed","id":"3(@)"}
86
+ {"time":"2026-08-14T23:11:19.670473153Z","level":"INFO","msg":"connection: ManageConnectionData: new connection created","id":"4(@)"}
87
+ {"time":"2026-08-14T23:11:19.883841265Z","level":"INFO","msg":"handleInformInit: received","streamId":"cvkq49dn","id":"4(@)"}
88
+ {"time":"2026-08-14T23:11:20.144645028Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"cvkq49dn","id":"4(@)"}
zain/Activation/wandb/run-20260814_231119-cvkq49dn/logs/debug-internal.log ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ {"time":"2026-08-14T23:11:19.883985323Z","level":"INFO","msg":"wandb-core"}
2
+ {"time":"2026-08-14T23:11:19.88411539Z","level":"INFO","msg":"stream: starting","core version":"0.28.1"}
3
+ {"time":"2026-08-14T23:11:20.144460721Z","level":"INFO","msg":"stream: created new stream","id":"cvkq49dn"}
4
+ {"time":"2026-08-14T23:11:20.14455917Z","level":"INFO","msg":"handler: started"}
5
+ {"time":"2026-08-14T23:11:20.144636487Z","level":"INFO","msg":"stream: started"}
6
+ {"time":"2026-08-14T23:11:20.144648599Z","level":"INFO","msg":"writer: started","stream_id":"cvkq49dn"}
7
+ {"time":"2026-08-14T23:11:20.144706894Z","level":"INFO","msg":"sender: started"}
8
+ {"time":"2026-08-14T23:11:20.584339046Z","level":"INFO","msg":"filestream: sending request","total_files":1,"console_offset":0,"console_lines":1}
9
+ {"time":"2026-08-14T23:11:20.683540169Z","level":"INFO","msg":"filestream: request sent","status":"200 OK"}
zain/Activation/wandb/run-20260814_231119-cvkq49dn/logs/debug.log ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_setup.py:_flush():81] Current SDK version is 0.28.1
2
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_setup.py:_flush():81] Configure stats pid to 2164011
3
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_setup.py:_flush():81] Loading settings from environment variables
4
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_init.py:setup_run_log_directory():729] Logging user logs to /mnt/data/zainulabideen/zain-exp/notebooks/zain/Activation/wandb/run-20260814_231119-cvkq49dn/logs/debug.log
5
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_init.py:setup_run_log_directory():730] Logging internal logs to /mnt/data/zainulabideen/zain-exp/notebooks/zain/Activation/wandb/run-20260814_231119-cvkq49dn/logs/debug-internal.log
6
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_init.py:init():772] calling init triggers
7
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_init.py:init():777] wandb.init called with sweep_config: {}
8
+ config: {'_wandb': {}}
9
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_init.py:init():820] starting backend
10
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_init.py:init():826] Connected to an existing wandb-core service via WANDB_SERVICE
11
+ 2026-08-14 23:11:19,882 INFO MainThread:2164011 [wandb_init.py:init():835] sending inform_init request
12
+ 2026-08-14 23:11:20,145 INFO MainThread:2164011 [wandb_init.py:init():840] backend started and connected
13
+ 2026-08-14 23:11:20,148 INFO MainThread:2164011 [wandb_init.py:init():910] updated telemetry
14
+ 2026-08-14 23:11:20,155 INFO MainThread:2164011 [wandb_init.py:init():933] communicating run to backend with 90.0 second timeout
15
+ 2026-08-14 23:11:20,497 INFO MainThread:2164011 [wandb_init.py:init():978] starting run threads in backend
16
+ 2026-08-14 23:11:20,570 INFO MainThread:2164011 [wandb_run.py:_console_start():2621] atexit reg
17
+ 2026-08-14 23:11:20,570 INFO MainThread:2164011 [wandb_run.py:_redirect():2471] redirect: wrap_raw
18
+ 2026-08-14 23:11:20,570 INFO MainThread:2164011 [wandb_run.py:_redirect():2540] Wrapping output streams.
19
+ 2026-08-14 23:11:20,570 INFO MainThread:2164011 [wandb_run.py:_redirect():2563] Redirects installed.
20
+ 2026-08-14 23:11:20,573 INFO MainThread:2164011 [wandb_init.py:init():1016] run started, returning control to user process
21
+ 2026-08-14 23:11:20,574 INFO MainThread:2164011 [wandb_run.py:_config_callback():1346] config_cb None None {'transformers_version': '5.16.0.dev0', 'architectures': None, 'output_hidden_states': False, 'return_dict': True, 'dtype': None, 'chunk_size_feed_forward': 0, 'is_encoder_decoder': False, 'id2label': {0: 'LABEL_0', 1: 'LABEL_1'}, 'label2id': {'LABEL_0': 0, 'LABEL_1': 1}, 'problem_type': None, 'vocab_size': 4096, 'hidden_size': 128, 'intermediate_size': 256, 'num_hidden_layers': 100, 'num_attention_heads': 4, 'num_key_value_heads': 4, 'hidden_act': 'silu', 'max_position_embeddings': 512, 'initializer_range': 0.02, 'rms_norm_eps': 1e-06, 'use_cache': False, 'pad_token_id': 0, 'bos_token_id': 1, 'eos_token_id': 2, 'pretraining_tp': 1, 'tie_word_embeddings': True, 'rope_parameters': {'rope_theta': 10000.0, 'rope_type': 'default'}, 'attention_bias': False, 'attention_dropout': 0.0, 'mlp_bias': False, 'head_dim': 32, '_name_or_path': '', 'tokenizer_name': 'w-ahmad/tiny-stories-tokenizer', 'mlp_type': 'glu', 'activation': 'waleed10', 'waleed_beta': 2, 'model_type': 'tiny_llama', 'output_attentions': False, 'output_dir': 'out/glu-waleed10-100L_trash_run', 'per_device_train_batch_size': 64, 'num_train_epochs': 1, 'max_steps': 3500, 'learning_rate': 9e-05, 'lr_scheduler_type': 'constant_with_warmup', 'lr_scheduler_kwargs': None, 'warmup_steps': 50, 'optim': 'adamw_torch_fused', 'optim_args': None, 'weight_decay': 0.0, 'adam_beta1': 0.9, 'adam_beta2': 0.999, 'adam_epsilon': 1e-08, 'optim_target_modules': None, 'gradient_accumulation_steps': 1, 'average_tokens_across_devices': True, 'max_grad_norm': 1.0, 'label_smoothing_factor': 0.0, 'bf16': True, 'fp16': False, 'bf16_full_eval': False, 'fp16_full_eval': False, 'tf32': None, 'gradient_checkpointing': False, 'gradient_checkpointing_kwargs': None, 'torch_compile': False, 'torch_compile_backend': None, 'torch_compile_mode': None, 'use_liger_kernel': False, 'liger_kernel_config': None, 'neftune_noise_alpha': None, 'torch_empty_cache_steps': None, 'auto_find_batch_size': False, 'logging_strategy': 'steps', 'logging_steps': 20, 'logging_first_step': False, 'log_on_each_node': True, 'logging_nan_inf_filter': True, 'include_num_input_tokens_seen': 'no', 'log_level': 'passive', 'log_level_replica': 'warning', 'disable_tqdm': False, 'report_to': ['wandb'], 'run_name': 'LM-glu-waleed10-100L-16.9M-20260814-231118', 'project': 'huggingface', 'trackio_space_id': None, 'trackio_bucket_id': None, 'trackio_static_space_id': None, 'eval_strategy': 'steps', 'eval_steps': 2498, 'eval_delay': 0, 'per_device_eval_batch_size': 1024, 'prediction_loss_only': False, 'eval_on_start': False, 'eval_do_concat_batches': True, 'eval_use_gather_object': False, 'eval_accumulation_steps': None, 'include_for_metrics': [], 'batch_eval_metrics': False, 'save_only_model': False, 'save_strategy': 'steps', 'save_steps': 50, 'save_on_each_node': False, 'save_total_limit': None, 'enable_jit_checkpoint': False, 'push_to_hub': False, 'hub_token': '<HUB_TOKEN>', 'hub_private_repo': None, 'hub_model_id': 'w-ahmad/6L-glu-waleed10-100L', 'hub_strategy': 'every_save', 'hub_always_push': False, 'hub_revision': None, 'load_best_model_at_end': False, 'metric_for_best_model': None, 'greater_is_better': None, 'ignore_data_skip': False, 'restore_callback_states_from_checkpoint': False, 'full_determinism': False, 'seed': 42, 'data_seed': 42, 'use_cpu': False, 'accelerator_config': {'split_batches': False, 'dispatch_batches': None, 'even_batches': True, 'use_seedable_sampler': True, 'non_blocking': False, 'gradient_accumulation_kwargs': None}, 'parallelism_config': None, 'dataloader_drop_last': False, 'dataloader_num_workers': 0, 'dataloader_pin_memory': True, 'dataloader_persistent_workers': False, 'dataloader_prefetch_factor': None, 'dataloader_multiprocessing_context': None, 'dataloader_in_order': True, 'remove_unused_columns': False, 'label_names': None, 'train_sampling_strategy': 'random', 'length_column_name': 'length', 'ddp_find_unused_parameters': None, 'ddp_bucket_cap_mb': None, 'ddp_broadcast_buffers': None, 'ddp_static_graph': None, 'ddp_backend': None, 'ddp_timeout': 1800, 'fsdp': None, 'fsdp_config': None, 'deepspeed': None, 'debug': [], 'skip_memory_metrics': True, 'do_train': False, 'do_eval': True, 'do_predict': False, 'resume_from_checkpoint': None, 'local_rank': -1}
22
+ 2026-08-14 23:11:20,578 INFO MainThread:2164011 [wandb_config.py:__setitem__():155] [no run ID] config set model/num_parameters = 16934016 - <bound method Run._config_callback of <wandb.sdk.wandb_run.Run object at 0x14c6ee183d10>>
23
+ 2026-08-14 23:11:20,578 INFO MainThread:2164011 [wandb_run.py:_config_callback():1346] config_cb model/num_parameters 16934016 None
zain/Activation/wandb/run-20260814_231119-cvkq49dn/run-cvkq49dn.wandb ADDED
Binary file (7 Bytes). View file