w-ahmad commited on
Commit
82dc38f
·
verified ·
1 Parent(s): 9b4c79a

Auto upload zain 2026-08-14T19:47:24.937839

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +1 -0
  2. zain/Activation/__pycache__/exp.cpython-311.pyc +0 -0
  3. zain/Activation/exp.py +227 -408
  4. zain/Activation/out/glu-relu-100L_run/checkpoint-100/config.json +36 -0
  5. zain/Activation/out/glu-relu-100L_run/checkpoint-100/model.safetensors +3 -0
  6. zain/Activation/out/glu-relu-100L_run/checkpoint-100/optimizer.pt +3 -0
  7. zain/Activation/out/glu-relu-100L_run/checkpoint-100/rng_state.pth +3 -0
  8. zain/Activation/out/glu-relu-100L_run/checkpoint-100/scheduler.pt +3 -0
  9. zain/Activation/out/glu-relu-100L_run/checkpoint-100/tokenizer.json +0 -0
  10. zain/Activation/out/glu-relu-100L_run/checkpoint-100/tokenizer_config.json +13 -0
  11. zain/Activation/out/glu-relu-100L_run/checkpoint-100/trainer_state.json +69 -0
  12. zain/Activation/out/glu-relu-100L_run/checkpoint-100/training_args.bin +3 -0
  13. zain/Activation/out/glu-relu-100L_run/checkpoint-1000/config.json +36 -0
  14. zain/Activation/out/glu-relu-100L_run/checkpoint-1000/model.safetensors +3 -0
  15. zain/Activation/out/glu-relu-100L_run/checkpoint-1000/optimizer.pt +3 -0
  16. zain/Activation/out/glu-relu-100L_run/checkpoint-1000/rng_state.pth +3 -0
  17. zain/Activation/out/glu-relu-100L_run/checkpoint-1000/scheduler.pt +3 -0
  18. zain/Activation/out/glu-relu-100L_run/checkpoint-1000/tokenizer.json +0 -0
  19. zain/Activation/out/glu-relu-100L_run/checkpoint-1000/tokenizer_config.json +13 -0
  20. zain/Activation/out/glu-relu-100L_run/checkpoint-1000/trainer_state.json +384 -0
  21. zain/Activation/out/glu-relu-100L_run/checkpoint-1000/training_args.bin +3 -0
  22. zain/Activation/out/glu-relu-100L_run/checkpoint-1100/config.json +36 -0
  23. zain/Activation/out/glu-relu-100L_run/checkpoint-1100/model.safetensors +3 -0
  24. zain/Activation/out/glu-relu-100L_run/checkpoint-1100/optimizer.pt +3 -0
  25. zain/Activation/out/glu-relu-100L_run/checkpoint-1100/rng_state.pth +3 -0
  26. zain/Activation/out/glu-relu-100L_run/checkpoint-1100/scheduler.pt +3 -0
  27. zain/Activation/out/glu-relu-100L_run/checkpoint-1100/tokenizer.json +0 -0
  28. zain/Activation/out/glu-relu-100L_run/checkpoint-1100/tokenizer_config.json +13 -0
  29. zain/Activation/out/glu-relu-100L_run/checkpoint-1100/trainer_state.json +419 -0
  30. zain/Activation/out/glu-relu-100L_run/checkpoint-1100/training_args.bin +3 -0
  31. zain/Activation/out/glu-relu-100L_run/checkpoint-1200/config.json +36 -0
  32. zain/Activation/out/glu-relu-100L_run/checkpoint-1200/model.safetensors +3 -0
  33. zain/Activation/out/glu-relu-100L_run/checkpoint-1200/optimizer.pt +3 -0
  34. zain/Activation/out/glu-relu-100L_run/checkpoint-1200/rng_state.pth +3 -0
  35. zain/Activation/out/glu-relu-100L_run/checkpoint-1200/scheduler.pt +3 -0
  36. zain/Activation/out/glu-relu-100L_run/checkpoint-1200/tokenizer.json +0 -0
  37. zain/Activation/out/glu-relu-100L_run/checkpoint-1200/tokenizer_config.json +13 -0
  38. zain/Activation/out/glu-relu-100L_run/checkpoint-1200/trainer_state.json +454 -0
  39. zain/Activation/out/glu-relu-100L_run/checkpoint-1200/training_args.bin +3 -0
  40. zain/Activation/out/glu-relu-100L_run/checkpoint-1300/config.json +36 -0
  41. zain/Activation/out/glu-relu-100L_run/checkpoint-1300/model.safetensors +3 -0
  42. zain/Activation/out/glu-relu-100L_run/checkpoint-1300/optimizer.pt +3 -0
  43. zain/Activation/out/glu-relu-100L_run/checkpoint-1300/rng_state.pth +3 -0
  44. zain/Activation/out/glu-relu-100L_run/checkpoint-1300/scheduler.pt +3 -0
  45. zain/Activation/out/glu-relu-100L_run/checkpoint-1300/tokenizer.json +0 -0
  46. zain/Activation/out/glu-relu-100L_run/checkpoint-1300/tokenizer_config.json +13 -0
  47. zain/Activation/out/glu-relu-100L_run/checkpoint-1300/trainer_state.json +489 -0
  48. zain/Activation/out/glu-relu-100L_run/checkpoint-1300/training_args.bin +3 -0
  49. zain/Activation/out/glu-relu-100L_run/checkpoint-1400/config.json +36 -0
  50. zain/Activation/out/glu-relu-100L_run/checkpoint-1400/model.safetensors +3 -0
.gitattributes CHANGED
@@ -41,3 +41,4 @@ zain/Activation/wandb/run-20260812_225525-hf77resg/run-hf77resg.wandb filter=lfs
41
  zain/Activation/wandb/run-20260812_230915-vvjf0upl/run-vvjf0upl.wandb filter=lfs diff=lfs merge=lfs -text
42
  zain/Activation/wandb/run-20260812_232326-8u0k1nti/run-8u0k1nti.wandb filter=lfs diff=lfs merge=lfs -text
43
  zain/Activation/wandb/run-20260812_233704-tpo5j00e/run-tpo5j00e.wandb filter=lfs diff=lfs merge=lfs -text
 
 
41
  zain/Activation/wandb/run-20260812_230915-vvjf0upl/run-vvjf0upl.wandb filter=lfs diff=lfs merge=lfs -text
42
  zain/Activation/wandb/run-20260812_232326-8u0k1nti/run-8u0k1nti.wandb filter=lfs diff=lfs merge=lfs -text
43
  zain/Activation/wandb/run-20260812_233704-tpo5j00e/run-tpo5j00e.wandb filter=lfs diff=lfs merge=lfs -text
44
+ zain/Activation/wandb/run-20260814_193907-kk1xbqht/run-kk1xbqht.wandb filter=lfs diff=lfs merge=lfs -text
zain/Activation/__pycache__/exp.cpython-311.pyc CHANGED
Binary files a/zain/Activation/__pycache__/exp.cpython-311.pyc and b/zain/Activation/__pycache__/exp.cpython-311.pyc differ
 
zain/Activation/exp.py CHANGED
@@ -1,8 +1,6 @@
1
- """
2
- Tiny Llama GLU Lab Consolidated training library.
3
- One file: model definition, activation registry, stability monitoring,
4
- time tracking, dataset builder, and Trainer factory.
5
- """
6
 
7
  import math
8
  import os
@@ -32,6 +30,7 @@ from transformers.models.llama.modeling_llama import (
32
  )
33
  from transformers.modeling_outputs import CausalLMOutputWithPast
34
  from datasets import load_dataset
 
35
 
36
 
37
  # =============================================================================
@@ -39,7 +38,6 @@ from datasets import load_dataset
39
  # =============================================================================
40
 
41
  class GLUActivationRegistry:
42
- """Own every gating activation you test. Add new variants in one line."""
43
  _registry: Dict[str, Callable[[torch.Tensor], torch.Tensor]] = {}
44
 
45
  @classmethod
@@ -54,7 +52,6 @@ class GLUActivationRegistry:
54
  )
55
  return cls._registry[name]
56
 
57
-
58
  # Built-ins
59
  GLUActivationRegistry.register("silu", nn.functional.silu)
60
  GLUActivationRegistry.register("swish", nn.functional.silu)
@@ -64,8 +61,6 @@ GLUActivationRegistry.register("sigmoid", torch.sigmoid)
64
  GLUActivationRegistry.register("tanh", torch.tanh)
65
  GLUActivationRegistry.register("softplus", nn.functional.softplus)
66
  GLUActivationRegistry.register("linear", lambda x: x)
67
-
68
- # Custom
69
  GLUActivationRegistry.register("s10", lambda x: x * x * torch.sigmoid(x))
70
  GLUActivationRegistry.register("w1a", lambda x: x * torch.tanh(x))
71
 
@@ -75,13 +70,6 @@ GLUActivationRegistry.register("w1a", lambda x: x * torch.tanh(x))
75
  # =============================================================================
76
 
77
  class TinyLlamaConfig(LlamaConfig):
78
- """
79
- Exact Llama config plus two fields:
80
- - mlp_type: "glu" or "mlp" (standard Transformer MLP)
81
- - activation: name of the activation function to use inside the MLP block.
82
- - waleed_beta: β_cap for the waleed10 post‑clip (default 10.0)
83
- Enforces pure MHA by requiring num_key_value_heads == num_attention_heads.
84
- """
85
  model_type = "tiny_llama"
86
 
87
  def __init__(
@@ -103,20 +91,13 @@ class TinyLlamaConfig(LlamaConfig):
103
 
104
 
105
  # =============================================================================
106
- # 3. MODEL
107
  # =============================================================================
108
 
109
  class TinyLlamaMLP(nn.Module):
110
- """
111
- Unified MLP block supporting:
112
- - Standard GLU: down_proj( act(gate_proj(x)) * up_proj(x) )
113
- - Standard MLP: down_proj( act(up_proj(x)) )
114
- - Gated variants (situglu, waleed) with internal tanh scaling.
115
- - waleed10 / silu-waleed10 with post‑down‑proj clipping.
116
-
117
- All projections are always named the same way, ensuring consistent
118
- tensor logging regardless of activation.
119
- """
120
  def __init__(self, config: TinyLlamaConfig):
121
  super().__init__()
122
  self.hidden_size = config.hidden_size
@@ -125,27 +106,25 @@ class TinyLlamaMLP(nn.Module):
125
  self.activation_name = config.activation
126
  self.waleed_beta = getattr(config, "waleed_beta", 10.0)
127
 
128
- # Effective intermediate size (scaled for MLP)
129
  if self.mlp_type == "glu":
130
  effective_intermediate = self.intermediate_size
131
  elif self.mlp_type == "mlp":
132
  effective_intermediate = int(self.intermediate_size * 1.5)
133
- print(f"[MLP] Auto‑scaled intermediate_size from {self.intermediate_size} to {effective_intermediate} for parameter parity.")
134
  else:
135
  raise ValueError(f"Unknown mlp_type: {self.mlp_type}")
136
 
137
  self.effective_intermediate = effective_intermediate
138
 
139
- # Always define the projections (names consistent across variants)
140
  if self.mlp_type == "glu":
141
  self.gate_proj = nn.Linear(self.hidden_size, effective_intermediate, bias=False)
142
  self.up_proj = nn.Linear(self.hidden_size, effective_intermediate, bias=False)
143
- else: # mlp
144
  self.up_proj = nn.Linear(self.hidden_size, effective_intermediate, bias=False)
145
 
146
  self.down_proj = nn.Linear(effective_intermediate, self.hidden_size, bias=False)
147
 
148
- # Activation function (for standard variants)
149
  if self.mlp_type == "glu":
150
  if self.activation_name in ("situglu", "waleed", "situglu_low", "waleedglu_low"):
151
  self.act_fn = None
@@ -155,12 +134,9 @@ class TinyLlamaMLP(nn.Module):
155
  self.act_fn = GLUActivationRegistry.get("silu")
156
  else:
157
  self.act_fn = GLUActivationRegistry.get(self.activation_name)
158
- else: # mlp
159
  if self.activation_name in ("situglu", "waleed", "situglu_low", "waleedglu_low"):
160
- raise ValueError(
161
- f"Activation '{self.activation_name}' requires a gated architecture (GLU). "
162
- f"Please use mlp_type='glu'."
163
- )
164
  elif self.activation_name == "waleed10":
165
  self.act_fn = GLUActivationRegistry.get("linear")
166
  elif self.activation_name == "silu-waleed10":
@@ -168,7 +144,6 @@ class TinyLlamaMLP(nn.Module):
168
  else:
169
  self.act_fn = GLUActivationRegistry.get(self.activation_name)
170
 
171
- # Determine beta values for situglu/waleed variants
172
  if self.activation_name in ("situglu_low", "waleedglu_low"):
173
  self.beta1 = 2.5
174
  self.beta2 = 4.0
@@ -176,7 +151,6 @@ class TinyLlamaMLP(nn.Module):
176
  self.beta1 = 4.0
177
  self.beta2 = 25.0
178
 
179
- # Flags for special handling
180
  self.is_situglu = self.activation_name in ("situglu", "situglu_low")
181
  self.is_waleed = self.activation_name in ("waleed", "waleedglu_low")
182
  self.is_waleed10 = self.activation_name in ("waleed10", "silu-waleed10")
@@ -206,9 +180,17 @@ class TinyLlamaMLP(nn.Module):
206
  if self.is_waleed10:
207
  out = self.waleed_beta * torch.tanh(out / self.waleed_beta)
208
 
 
 
 
 
209
  return out
210
 
211
 
 
 
 
 
212
  class TinyLlamaDecoderLayer(nn.Module):
213
  def __init__(self, config: TinyLlamaConfig, layer_idx: int):
214
  super().__init__()
@@ -216,23 +198,12 @@ class TinyLlamaDecoderLayer(nn.Module):
216
  self.self_attn = LlamaAttention(config=config, layer_idx=layer_idx)
217
  self.mlp = TinyLlamaMLP(config)
218
  self.input_layernorm = LlamaRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
219
- self.post_attention_layernorm = LlamaRMSNorm(
220
- config.hidden_size, eps=config.rms_norm_eps
221
- )
222
- # Residual stream gateways (zero params, hookable)
223
  self.residual_pre_attn = nn.Identity()
224
  self.residual_post_attn = nn.Identity()
225
  self.residual_post_mlp = nn.Identity()
226
 
227
- def forward(
228
- self,
229
- hidden_states: torch.Tensor,
230
- attention_mask: Optional[torch.Tensor] = None,
231
- position_ids: Optional[torch.LongTensor] = None,
232
- position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
233
- **kwargs,
234
- ):
235
- # Attention sub-layer
236
  residual = hidden_states
237
  hidden_states = self.residual_pre_attn(hidden_states)
238
  hidden_states = self.input_layernorm(hidden_states)
@@ -245,7 +216,6 @@ class TinyLlamaDecoderLayer(nn.Module):
245
  hidden_states = residual + attn_out
246
  hidden_states = self.residual_post_attn(hidden_states)
247
 
248
- # MLP sub-layer
249
  residual = hidden_states
250
  hidden_states = self.post_attention_layernorm(hidden_states)
251
  hidden_states = self.mlp(hidden_states)
@@ -254,24 +224,14 @@ class TinyLlamaDecoderLayer(nn.Module):
254
  return (hidden_states,)
255
 
256
 
257
- # ----------------------------------------------------------------------------
258
- # ATTENTION MASK
259
- # ----------------------------------------------------------------------------
260
- def _build_causal_mask(
261
- attention_mask: Optional[torch.Tensor],
262
- seq_len: int,
263
- dtype: torch.dtype,
264
- device: torch.device,
265
- ) -> torch.Tensor:
266
  min_value = torch.finfo(dtype).min
267
  causal = torch.full((seq_len, seq_len), fill_value=min_value, dtype=dtype, device=device)
268
  causal = torch.triu(causal, diagonal=1)
269
  causal = causal[None, None, :, :]
270
-
271
  if attention_mask is None:
272
  batch_size = 1
273
  return causal.expand(batch_size, 1, seq_len, seq_len)
274
-
275
  batch_size = attention_mask.shape[0]
276
  causal = causal.expand(batch_size, 1, seq_len, seq_len).clone()
277
  padding = attention_mask[:, None, None, :].to(device) == 0
@@ -289,48 +249,27 @@ class TinyLlamaModel(LlamaPreTrainedModel):
289
  super().__init__(config)
290
  self.padding_idx = config.pad_token_id
291
  self.vocab_size = config.vocab_size
292
- self.embed_tokens = nn.Embedding(
293
- config.vocab_size, config.hidden_size, self.padding_idx
294
- )
295
- self.layers = nn.ModuleList(
296
- [TinyLlamaDecoderLayer(config, i) for i in range(config.num_hidden_layers)]
297
- )
298
  self.norm = LlamaRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
299
  self.rotary_emb = LlamaRotaryEmbedding(config=config)
300
  self.post_init()
301
 
302
- def forward(
303
- self,
304
- input_ids: Optional[torch.LongTensor] = None,
305
- attention_mask: Optional[torch.Tensor] = None,
306
- position_ids: Optional[torch.LongTensor] = None,
307
- inputs_embeds: Optional[torch.FloatTensor] = None,
308
- return_dict: Optional[bool] = None,
309
- **kwargs,
310
- ):
311
  global _MASK_PRINTED
312
  return_dict = return_dict if return_dict is not None else self.config.use_return_dict
313
  if inputs_embeds is None:
314
  inputs_embeds = self.embed_tokens(input_ids)
315
-
316
  if position_ids is None:
317
  seq_len = inputs_embeds.shape[1]
318
- position_ids = torch.arange(
319
- seq_len, device=inputs_embeds.device
320
- ).unsqueeze(0).expand(inputs_embeds.shape[0], -1)
321
-
322
  hidden_states = inputs_embeds
323
  position_embeddings = self.rotary_emb(hidden_states, position_ids)
324
-
325
  seq_len = hidden_states.shape[1]
326
- causal_mask = _build_causal_mask(
327
- attention_mask, seq_len, hidden_states.dtype, hidden_states.device
328
- )
329
-
330
  if not _MASK_PRINTED:
331
  print("[INFO] Causal mask (float with -inf) applied to all attention layers.")
332
  _MASK_PRINTED = True
333
-
334
  for decoder_layer in self.layers:
335
  layer_outputs = decoder_layer(
336
  hidden_states,
@@ -339,7 +278,6 @@ class TinyLlamaModel(LlamaPreTrainedModel):
339
  position_embeddings=position_embeddings,
340
  )
341
  hidden_states = layer_outputs[0]
342
-
343
  hidden_states = self.norm(hidden_states)
344
  if not return_dict:
345
  return (hidden_states,)
@@ -367,50 +305,25 @@ class TinyLlamaForCausalLM(LlamaPreTrainedModel):
367
  def get_output_embeddings(self):
368
  return self.lm_head
369
 
370
- def forward(
371
- self,
372
- input_ids: Optional[torch.LongTensor] = None,
373
- attention_mask: Optional[torch.Tensor] = None,
374
- position_ids: Optional[torch.LongTensor] = None,
375
- inputs_embeds: Optional[torch.FloatTensor] = None,
376
- labels: Optional[torch.LongTensor] = None,
377
- return_dict: Optional[bool] = None,
378
- **kwargs,
379
- ):
380
  return_dict = return_dict if return_dict is not None else self.config.use_return_dict
381
- outputs = self.model(
382
- input_ids=input_ids,
383
- attention_mask=attention_mask,
384
- position_ids=position_ids,
385
- inputs_embeds=inputs_embeds,
386
- return_dict=return_dict,
387
- )
388
  hidden_states = outputs["last_hidden_state"] if return_dict else outputs[0]
389
  logits = self.lm_head(hidden_states)
390
-
391
  loss = None
392
  if labels is not None:
393
  shift_logits = logits[..., :-1, :].contiguous()
394
  shift_labels = labels[..., 1:].contiguous()
395
  loss_fct = nn.CrossEntropyLoss()
396
- loss = loss_fct(
397
- shift_logits.view(-1, self.config.vocab_size), shift_labels.view(-1)
398
- )
399
-
400
  if not return_dict:
401
  output = (logits,) + outputs[1:]
402
  return (loss,) + output if loss is not None else output
403
- return CausalLMOutputWithPast(
404
- loss=loss,
405
- logits=logits,
406
- past_key_values=None,
407
- hidden_states=None,
408
- attentions=None,
409
- )
410
 
411
- def prepare_inputs_for_generation(
412
- self, input_ids, past_key_values=None, attention_mask=None, **kwargs
413
- ):
414
  if past_key_values:
415
  input_ids = input_ids[:, -1:]
416
  position_ids = kwargs.get("position_ids")
@@ -428,57 +341,141 @@ class TinyLlamaForCausalLM(LlamaPreTrainedModel):
428
 
429
 
430
  # =============================================================================
431
- # 4. MONITORING ENGINE
432
  # =============================================================================
433
 
434
- class StatsEngine:
435
- """Compute unified signature for any tensor."""
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
436
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
437
  @staticmethod
438
- def compute(
439
- tensor: torch.Tensor, user_limit: float, dtype_ratio: float
440
- ) -> Dict[str, float]:
441
  with torch.no_grad():
442
  abs_t = tensor.abs()
443
  dtype_info = torch.finfo(tensor.dtype)
444
- dtype_limit = (
445
- dtype_ratio * dtype_info.max
446
- if not torch.isinf(torch.tensor(dtype_info.max))
447
- else float("inf")
448
- )
449
-
450
- t_min = tensor.min().item()
451
- t_max = tensor.max().item()
452
-
453
  return {
454
  "norm": tensor.norm(2).item(),
455
  "mean": tensor.mean().item(),
456
  "std": tensor.std().item(),
457
  "max_abs": abs_t.max().item(),
458
- "frac_near_dtype_limit": (
459
- (abs_t > dtype_limit).float().mean().item()
460
- if not math.isinf(dtype_limit)
461
- else 0.0
462
- ),
463
  "frac_near_user_limit": (abs_t > user_limit).float().mean().item(),
464
- "min": t_min,
465
- "max": t_max,
466
- "range": t_max - t_min,
467
  }
468
 
469
 
470
  class StepAccumulator:
471
- """Stores per-tensor entries, aggregates to layer/global scope."""
472
-
473
  def __init__(self):
474
  self.tensors: Dict[str, Dict[str, float]] = {}
475
 
476
  def add(self, name: str, numel: int, stats: Dict[str, float]):
477
  new_entry = {"numel": numel, **stats}
478
  existing = self.tensors.get(name)
479
- self.tensors[name] = (
480
- new_entry if existing is None else self._merge_entry(existing, new_entry)
481
- )
482
 
483
  @staticmethod
484
  def _merge_entry(a: Dict[str, float], b: Dict[str, float]) -> Dict[str, float]:
@@ -488,21 +485,12 @@ class StepAccumulator:
488
  norm = math.sqrt(a["norm"] ** 2 + b["norm"] ** 2)
489
  max_abs = max(a["max_abs"], b["max_abs"])
490
  mean = (a["mean"] * a["numel"] + b["mean"] * b["numel"]) / total_n
491
- ex2 = (
492
- a["numel"] * (a["std"] ** 2 + a["mean"] ** 2)
493
- + b["numel"] * (b["std"] ** 2 + b["mean"] ** 2)
494
- ) / total_n
495
  std = math.sqrt(max(0.0, ex2 - mean ** 2))
496
- frac_dtype = (
497
- a["frac_near_dtype_limit"] * a["numel"] + b["frac_near_dtype_limit"] * b["numel"]
498
- ) / total_n
499
- frac_user = (
500
- a["frac_near_user_limit"] * a["numel"] + b["frac_near_user_limit"] * b["numel"]
501
- ) / total_n
502
-
503
  t_min = min(a.get("min", float("inf")), b.get("min", float("inf")))
504
  t_max = max(a.get("max", float("-inf")), b.get("max", float("-inf")))
505
-
506
  return {
507
  "numel": total_n,
508
  "norm": norm,
@@ -524,27 +512,15 @@ class StepAccumulator:
524
  return {}
525
  numels = [e["numel"] for e in entries.values()]
526
  total_n = sum(numels)
527
-
528
  norm = math.sqrt(sum(e["norm"] ** 2 for e in entries.values()))
529
  max_abs = max(e["max_abs"] for e in entries.values())
530
  mean = sum(e["mean"] * e["numel"] for e in entries.values()) / total_n
531
- ex2 = (
532
- sum(e["numel"] * (e["std"] ** 2 + e["mean"] ** 2) for e in entries.values())
533
- / total_n
534
- )
535
  std = math.sqrt(max(0.0, ex2 - mean ** 2))
536
- frac_dtype = (
537
- sum(e["frac_near_dtype_limit"] * e["numel"] for e in entries.values())
538
- / total_n
539
- )
540
- frac_user = (
541
- sum(e["frac_near_user_limit"] * e["numel"] for e in entries.values())
542
- / total_n
543
- )
544
-
545
  t_min = min(e.get("min", float("inf")) for e in entries.values())
546
  t_max = max(e.get("max", float("-inf")) for e in entries.values())
547
-
548
  return {
549
  "norm": norm,
550
  "mean": mean,
@@ -561,27 +537,17 @@ class StepAccumulator:
561
  return self._aggregate(self.tensors)
562
 
563
  def get_layer_stats(self, layer_prefix: str) -> Dict[str, float]:
564
- entries = {
565
- k: v for k, v in self.tensors.items() if k.startswith(layer_prefix + ".")
566
- }
567
  return self._aggregate(entries)
568
 
569
 
570
  class HookRegistry:
571
- """Attach and throttle forward/backward hooks."""
572
-
573
  def __init__(self, model: nn.Module):
574
  self.model = model
575
- self.handles: List[torch.utils.hooks.RemovableHandle] = []
576
  self.active = False
577
 
578
- def attach_forward(
579
- self,
580
- module_patterns: List[str],
581
- accumulator: StepAccumulator,
582
- user_limit: float,
583
- dtype_ratio: float,
584
- ):
585
  for name, module in self.model.named_modules():
586
  if any(re.search(p, name) for p in module_patterns):
587
  h = module.register_forward_hook(
@@ -589,51 +555,35 @@ class HookRegistry:
589
  )
590
  self.handles.append(h)
591
 
592
- def attach_backward(
593
- self, accumulator: StepAccumulator, user_limit: float, dtype_ratio: float
594
- ):
595
  for name, param in self.model.named_parameters():
596
- if param.requires_grad:
597
- h = param.register_hook(
598
- self._make_backward_hook(
599
- f"grad.{name}", accumulator, user_limit, dtype_ratio
600
- )
601
- )
602
- self.handles.append(h)
 
603
 
604
- def _make_forward_hook(
605
- self, module_name: str, accumulator: StepAccumulator, user_limit: float, dtype_ratio: float
606
- ):
607
  def hook(module, inp, out):
608
  if not self.active:
609
  return
610
  if isinstance(out, dict):
611
- out_dict = out
612
- out = out_dict.get("last_hidden_state")
613
- if out is None:
614
- out = next(
615
- (v for v in out_dict.values() if torch.is_tensor(v)), None
616
- )
617
- elif isinstance(out, (tuple, list)):
618
- out = out[0] if len(out) > 0 else None
619
-
620
  if not torch.is_tensor(out):
621
  return
622
-
623
  stats = StatsEngine.compute(out.detach(), user_limit, dtype_ratio)
624
  accumulator.add(f"act.{module_name}", out.numel(), stats)
625
-
626
  return hook
627
 
628
- def _make_backward_hook(
629
- self, param_name: str, accumulator: StepAccumulator, user_limit: float, dtype_ratio: float
630
- ):
631
  def hook(grad):
632
  if not self.active:
633
  return
634
  stats = StatsEngine.compute(grad.detach(), user_limit, dtype_ratio)
635
  accumulator.add(param_name, grad.numel(), stats)
636
-
637
  return hook
638
 
639
  def set_active(self, active: bool):
@@ -646,13 +596,12 @@ class HookRegistry:
646
 
647
 
648
  class StabilityMonitorCallback(TrainerCallback):
649
- """Full stability instrumentation: grad / param / act statistics."""
650
-
651
  def __init__(
652
  self,
653
  model: nn.Module,
654
  monitor_every_n_steps: int = 10,
655
  module_patterns: Optional[List[str]] = None,
 
656
  user_limits: Optional[Dict[str, float]] = None,
657
  dtype_proximity_ratio: float = 0.9,
658
  log_scope: Optional[Dict[str, bool]] = None,
@@ -660,14 +609,11 @@ class StabilityMonitorCallback(TrainerCallback):
660
  ):
661
  self.model = model
662
  self.monitor_every_n_steps = monitor_every_n_steps
663
- self.module_patterns = module_patterns or [".*mlp.*", ".*self_attn.*", ".*residual.*"]
 
664
  self.user_limits = user_limits or {"grad": 1.0, "param": 100.0, "act": 50.0}
665
  self.dtype_ratio = dtype_proximity_ratio
666
- self.log_scope = log_scope or {
667
- "global": True,
668
- "per_layer": True,
669
- "per_tensor": False,
670
- }
671
  self.monitor_during_eval = monitor_during_eval
672
 
673
  self.accumulator = StepAccumulator()
@@ -679,12 +625,14 @@ class StabilityMonitorCallback(TrainerCallback):
679
  self.dtype_ratio,
680
  )
681
  self.hooks.attach_backward(
682
- self.accumulator, self.user_limits["grad"], self.dtype_ratio
 
 
 
683
  )
 
684
 
685
- self.pending_metrics: Optional[Dict[str, float]] = None
686
-
687
- def _should_monitor(self, state) -> bool:
688
  return state.global_step % self.monitor_every_n_steps == 0
689
 
690
  def on_step_begin(self, args, state, control, **kwargs):
@@ -696,16 +644,14 @@ class StabilityMonitorCallback(TrainerCallback):
696
  if not self.hooks.active:
697
  return
698
  for name, param in self.model.named_parameters():
699
- stats = StatsEngine.compute(
700
- param.data, self.user_limits["param"], self.dtype_ratio
701
- )
702
  self.accumulator.add(f"param.{name}", param.numel(), stats)
703
-
704
  self.hooks.set_active(False)
705
  self.pending_metrics = self._build_metrics()
706
 
707
- @staticmethod
708
- def _kind_of(name: str) -> str:
709
  if name.startswith("act."):
710
  return "act"
711
  if name.startswith("grad."):
@@ -714,8 +660,7 @@ class StabilityMonitorCallback(TrainerCallback):
714
  return "param"
715
  return "other"
716
 
717
- @staticmethod
718
- def _strip_kind(name: str) -> str:
719
  if name.startswith("act."):
720
  return name[4:]
721
  if name.startswith(("grad.", "param.")):
@@ -723,10 +668,9 @@ class StabilityMonitorCallback(TrainerCallback):
723
  return name
724
 
725
  def _build_metrics(self, scope: str = "train") -> Dict[str, float]:
726
- metrics: Dict[str, float] = {}
727
-
728
  if self.log_scope.get("global", True):
729
- by_kind: Dict[str, Dict[str, Dict[str, float]]] = {}
730
  for k, v in self.accumulator.tensors.items():
731
  by_kind.setdefault(self._kind_of(k), {})[k] = v
732
  for kind, entries in by_kind.items():
@@ -744,7 +688,7 @@ class StabilityMonitorCallback(TrainerCallback):
744
  prefix = ".".join(parts[: i + 2])
745
  layer_prefixes.add(prefix)
746
  for prefix in layer_prefixes:
747
- by_kind: Dict[str, Dict[str, Dict[str, float]]] = {}
748
  for k, v in self.accumulator.tensors.items():
749
  clean = self._strip_kind(k)
750
  if clean.startswith(prefix + ".") or clean == prefix:
@@ -764,12 +708,17 @@ class StabilityMonitorCallback(TrainerCallback):
764
  if kk == "numel":
765
  continue
766
  metrics[f"{scope}/tensor_{safe}/{kk}"] = vv
767
-
768
  return metrics
769
 
770
  def on_log(self, args, state, control, logs=None, **kwargs):
771
  if logs is not None and self.pending_metrics is not None:
772
  logs.update(self.pending_metrics)
 
 
 
 
 
 
773
  self.pending_metrics = None
774
 
775
  def on_prediction_step(self, args, state, control, **kwargs):
@@ -779,9 +728,9 @@ class StabilityMonitorCallback(TrainerCallback):
779
  self.accumulator.clear()
780
  self.hooks.set_active(True)
781
  for name, param in self.model.named_parameters():
782
- stats = StatsEngine.compute(
783
- param.data, self.user_limits["param"], self.dtype_ratio
784
- )
785
  self.accumulator.add(f"param.{name}", param.numel(), stats)
786
  self.pending_metrics = self._build_metrics(scope="eval")
787
 
@@ -790,14 +739,16 @@ class StabilityMonitorCallback(TrainerCallback):
790
  self.accumulator.clear()
791
 
792
 
793
- class TimeTrackerCallback(TrainerCallback):
794
- """Precise training & eval timing with remaining-time estimates."""
 
795
 
 
796
  def __init__(self):
797
- self.step_start: Optional[float] = None
798
- self.epoch_start: Optional[float] = None
799
  self.total_train_time = 0.0
800
- self.step_times: List[float] = []
801
 
802
  def on_epoch_begin(self, args, state, control, **kwargs):
803
  self.epoch_start = time.perf_counter()
@@ -826,13 +777,12 @@ class TimeTrackerCallback(TrainerCallback):
826
  remaining = (state.max_steps - state.global_step) * avg
827
  logs["train/estimated_remaining_minutes"] = remaining / 60.0
828
 
829
- def on_evaluate(self, args, state, control, metrics=None, **kwargs):
830
- pass
831
 
 
 
 
832
 
833
  class MetricsLoggerCallback(TrainerCallback):
834
- """Persist every logged dict as JSONL in the output dir."""
835
-
836
  def __init__(self, output_dir: str):
837
  self.output_dir = Path(output_dir)
838
  self.output_dir.mkdir(parents=True, exist_ok=True)
@@ -841,113 +791,13 @@ class MetricsLoggerCallback(TrainerCallback):
841
  def on_log(self, args, state, control, logs=None, **kwargs):
842
  if logs is None:
843
  return
844
- entry = {
845
- "step": state.global_step,
846
- "epoch": state.epoch,
847
- "timestamp": time.time(),
848
- **logs,
849
- }
850
  with open(self.log_file, "a") as f:
851
  f.write(json.dumps(entry, default=str) + "\n")
852
 
853
 
854
  # =============================================================================
855
- # CONTAMINATION CALLBACK
856
- # =============================================================================
857
-
858
- class ContaminationCallback(TrainerCallback):
859
- """
860
- Intentionally corrupt input_ids and labels for a window of steps.
861
- Modes: "shift" (+1 mod vocab), "random" (uniform random IDs).
862
- """
863
- def __init__(
864
- self,
865
- vocab_size: int,
866
- enabled: bool = False,
867
- start_step: int = 0,
868
- duration_steps: int = 0,
869
- mode: str = "shift",
870
- fraction: float = 1.0,
871
- seed: Optional[int] = None,
872
- ):
873
- self.vocab_size = vocab_size
874
- self.enabled = enabled
875
- self.start_step = start_step
876
- self.duration_steps = duration_steps
877
- self.mode = mode
878
- self.fraction = fraction
879
- self.seed = seed
880
- self.generator = torch.Generator()
881
- if seed is not None:
882
- self.generator.manual_seed(seed)
883
- self._active = False
884
-
885
- def _should_corrupt(self, state) -> bool:
886
- if not self.enabled:
887
- return False
888
- step = state.global_step
889
- return self.start_step <= step < self.start_step + self.duration_steps
890
-
891
- def on_step_begin(self, args, state, control, **kwargs):
892
- if not self._should_corrupt(state):
893
- return
894
- batch = kwargs.get("inputs")
895
- if batch is None or not isinstance(batch, dict):
896
- return
897
- input_ids = batch.get("input_ids")
898
- labels = batch.get("labels")
899
- attention_mask = batch.get("attention_mask")
900
- if input_ids is None or labels is None:
901
- return
902
-
903
- device = input_ids.device
904
- batch_size, seq_len = input_ids.shape
905
-
906
- if self.fraction >= 1.0:
907
- corrupt_mask = torch.ones((batch_size, seq_len), dtype=torch.bool, device=device)
908
- elif self.fraction <= 0.0:
909
- return
910
- else:
911
- rand = torch.rand((batch_size, seq_len), generator=self.generator, device=device)
912
- corrupt_mask = rand < self.fraction
913
-
914
- if attention_mask is not None:
915
- padding_mask = attention_mask == 0
916
- corrupt_mask = corrupt_mask & (~padding_mask)
917
-
918
- if not corrupt_mask.any():
919
- return
920
-
921
- if self.mode == "shift":
922
- input_ids_corrupt = (input_ids + 1) % self.vocab_size
923
- labels_corrupt = (labels + 1) % self.vocab_size
924
- elif self.mode == "random":
925
- random_ids = torch.randint(
926
- 0, self.vocab_size, input_ids.shape,
927
- generator=self.generator, device=device
928
- )
929
- input_ids_corrupt = random_ids
930
- labels_corrupt = random_ids
931
- else:
932
- raise ValueError(f"Unknown contamination mode: {self.mode}")
933
-
934
- input_ids.masked_scatter_(corrupt_mask, input_ids_corrupt[corrupt_mask])
935
- labels.masked_scatter_(corrupt_mask, labels_corrupt[corrupt_mask])
936
- batch["input_ids"] = input_ids
937
- batch["labels"] = labels
938
-
939
- if not self._active:
940
- print(f"[Contamination] Started at step {state.global_step} for {self.duration_steps} steps (mode={self.mode})")
941
- self._active = True
942
-
943
- def on_step_end(self, args, state, control, **kwargs):
944
- if self._active and state.global_step >= self.start_step + self.duration_steps:
945
- print(f"[Contamination] Ended at step {state.global_step}")
946
- self._active = False
947
-
948
-
949
- # =============================================================================
950
- # 5. DATA & TRAINER FACTORY
951
  # =============================================================================
952
 
953
  def build_dataset(
@@ -957,11 +807,9 @@ def build_dataset(
957
  dataset_name: str = "roneneldan/TinyStories",
958
  max_samples: Optional[int] = None,
959
  ):
960
- """Concatenate and chunk TinyStories for causal LM."""
961
  ds = load_dataset(dataset_name, split=split)
962
  if max_samples is not None and split == "train":
963
  ds = ds.select(range(min(max_samples, len(ds))))
964
- print(f"[Dataset] Using first {len(ds)} samples for training (max_samples={max_samples})")
965
 
966
  def tokenize(examples):
967
  out = tokenizer(examples["text"], add_special_tokens=False)
@@ -971,34 +819,20 @@ def build_dataset(
971
  out["attention_mask"] = [mask + [1] for mask in out["attention_mask"]]
972
  return out
973
 
974
- tokenized = ds.map(
975
- tokenize,
976
- batched=True,
977
- num_proc=4,
978
- remove_columns=ds.column_names,
979
- desc=f"Tokenizing {split}",
980
- )
981
 
982
  def group_texts(examples):
983
- concatenated = {
984
- k: list(chain.from_iterable(examples[k])) for k in examples.keys()
985
- }
986
  total_length = len(concatenated[list(examples.keys())[0]])
987
  total_length = (total_length // max_seq_len) * max_seq_len
988
  result = {
989
- k: [t[i : i + max_seq_len] for i in range(0, total_length, max_seq_len)]
990
  for k, t in concatenated.items()
991
  }
992
  result["labels"] = result["input_ids"].copy()
993
  return result
994
 
995
- return tokenized.map(
996
- group_texts,
997
- batched=True,
998
- batch_size=10000,
999
- num_proc=4,
1000
- desc=f"Chunking {split}",
1001
- )
1002
 
1003
 
1004
  def create_trainer(
@@ -1008,15 +842,17 @@ def create_trainer(
1008
  train_dataset,
1009
  eval_dataset=None,
1010
  ):
1011
- """Assemble HF Trainer with all custom callbacks."""
1012
  tc = config.get("training", {})
1013
  mc = config.get("monitor", {})
 
1014
 
1015
- run_name = tc.get("run_name", None)
 
 
1016
 
1017
  args = TrainingArguments(
1018
  output_dir=tc.get("output_dir", "./out"),
1019
- run_name=run_name,
1020
  num_train_epochs=tc.get("num_train_epochs", 3),
1021
  per_device_train_batch_size=tc.get("per_device_train_batch_size", 16),
1022
  per_device_eval_batch_size=tc.get("per_device_eval_batch_size", 16),
@@ -1046,27 +882,22 @@ def create_trainer(
1046
 
1047
  callbacks = [TimeTrackerCallback()]
1048
 
1049
- cc = config.get("contamination", {})
1050
- if cc.get("enabled", False):
1051
- vocab_size = model.config.vocab_size
1052
- callbacks.append(
1053
- ContaminationCallback(
1054
- vocab_size=vocab_size,
1055
- enabled=True,
1056
- start_step=cc.get("start_step", 0),
1057
- duration_steps=cc.get("duration_steps", 0),
1058
- mode=cc.get("mode", "shift"),
1059
- fraction=cc.get("fraction", 1.0),
1060
- seed=cc.get("seed", None),
1061
- )
1062
- )
1063
 
 
1064
  if mc.get("enabled", True):
 
1065
  callbacks.append(
1066
  StabilityMonitorCallback(
1067
  model=model,
1068
  monitor_every_n_steps=mc.get("monitor_every_n_steps", 10),
1069
- module_patterns=mc.get("module_patterns", [".*mlp.*", ".*self_attn.*", ".*residual.*"]),
 
1070
  user_limits=mc.get("user_limits", {"grad": 1.0, "param": 100.0, "act": 50.0}),
1071
  dtype_proximity_ratio=mc.get("dtype_proximity_ratio", 0.9),
1072
  log_scope=mc.get("log_scope", {"global": True, "per_layer": True, "per_tensor": False}),
@@ -1087,16 +918,4 @@ def create_trainer(
1087
  callbacks=callbacks,
1088
  )
1089
 
1090
- try:
1091
- from transformers.integrations import get_reporting_integration_callbacks
1092
- reporting_types = tuple(get_reporting_integration_callbacks(args.report_to))
1093
- except Exception:
1094
- reporting_types = ()
1095
-
1096
- if reporting_types:
1097
- handler = trainer.callback_handler
1098
- reporting_cbs = [cb for cb in handler.callbacks if isinstance(cb, reporting_types)]
1099
- other_cbs = [cb for cb in handler.callbacks if not isinstance(cb, reporting_types)]
1100
- handler.callbacks = other_cbs + reporting_cbs
1101
-
1102
  return trainer
 
1
+ # =====================================================================
2
+ # exp.py FULL FILE, ALL FIXES INCLUDED (hook requires_grad check)
3
+ # =====================================================================
 
 
4
 
5
  import math
6
  import os
 
30
  )
31
  from transformers.modeling_outputs import CausalLMOutputWithPast
32
  from datasets import load_dataset
33
+ from huggingface_hub import snapshot_download
34
 
35
 
36
  # =============================================================================
 
38
  # =============================================================================
39
 
40
  class GLUActivationRegistry:
 
41
  _registry: Dict[str, Callable[[torch.Tensor], torch.Tensor]] = {}
42
 
43
  @classmethod
 
52
  )
53
  return cls._registry[name]
54
 
 
55
  # Built-ins
56
  GLUActivationRegistry.register("silu", nn.functional.silu)
57
  GLUActivationRegistry.register("swish", nn.functional.silu)
 
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
 
 
70
  # =============================================================================
71
 
72
  class TinyLlamaConfig(LlamaConfig):
 
 
 
 
 
 
 
73
  model_type = "tiny_llama"
74
 
75
  def __init__(
 
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
 
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
 
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":
 
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
 
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")
 
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__()
 
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)
 
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)
 
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
 
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,
 
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,)
 
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")
 
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]:
 
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,
 
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,
 
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(
 
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):
 
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,
 
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()
 
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):
 
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."):
 
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.")):
 
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():
 
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:
 
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):
 
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
 
 
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()
 
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)
 
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(
 
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)
 
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(
 
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),
 
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}),
 
918
  callbacks=callbacks,
919
  )
920
 
 
 
 
 
 
 
 
 
 
 
 
 
921
  return trainer
zain/Activation/out/glu-relu-100L_run/checkpoint-100/config.json ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "activation": "relu",
3
+ "architectures": [
4
+ "TinyLlamaForCausalLM"
5
+ ],
6
+ "attention_bias": false,
7
+ "attention_dropout": 0.0,
8
+ "bos_token_id": 1,
9
+ "dtype": "bfloat16",
10
+ "eos_token_id": 2,
11
+ "head_dim": 32,
12
+ "hidden_act": "silu",
13
+ "hidden_size": 128,
14
+ "initializer_range": 0.02,
15
+ "intermediate_size": 256,
16
+ "max_position_embeddings": 512,
17
+ "mlp_bias": false,
18
+ "mlp_type": "glu",
19
+ "model_type": "tiny_llama",
20
+ "num_attention_heads": 4,
21
+ "num_hidden_layers": 100,
22
+ "num_key_value_heads": 4,
23
+ "pad_token_id": 0,
24
+ "pretraining_tp": 1,
25
+ "rms_norm_eps": 1e-06,
26
+ "rope_parameters": {
27
+ "rope_theta": 10000.0,
28
+ "rope_type": "default"
29
+ },
30
+ "tie_word_embeddings": true,
31
+ "tokenizer_name": "w-ahmad/tiny-stories-tokenizer",
32
+ "transformers_version": "5.16.0.dev0",
33
+ "use_cache": false,
34
+ "vocab_size": 4096,
35
+ "waleed_beta": 4
36
+ }
zain/Activation/out/glu-relu-100L_run/checkpoint-100/model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:a089d05e4a988d02ff08e00f5b97cf9ca33101dc26f5000abc24b4dfa7d8fb57
3
+ size 33967272
zain/Activation/out/glu-relu-100L_run/checkpoint-100/optimizer.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:54deacb71de71964a9b208737e1807784826d4de4553c32143fce95b417d6e16
3
+ size 68504996
zain/Activation/out/glu-relu-100L_run/checkpoint-100/rng_state.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:9d9cd6a0487226e5bd30d1846894c82af483733ab4381b75bae9c0745e05d405
3
+ size 14244
zain/Activation/out/glu-relu-100L_run/checkpoint-100/scheduler.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3001486ba00eb51d89ef078fe70e7e37535cd4e2311d7dd589266d045c9fb892
3
+ size 1064
zain/Activation/out/glu-relu-100L_run/checkpoint-100/tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
zain/Activation/out/glu-relu-100L_run/checkpoint-100/tokenizer_config.json ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_prefix_space": false,
3
+ "backend": "tokenizers",
4
+ "bos_token": "<|endoftext|>",
5
+ "eos_token": "<|endoftext|>",
6
+ "errors": "replace",
7
+ "is_local": false,
8
+ "local_files_only": false,
9
+ "model_max_length": 1024,
10
+ "pad_token": "<|endoftext|>",
11
+ "tokenizer_class": "GPT2Tokenizer",
12
+ "unk_token": "<|endoftext|>"
13
+ }
zain/Activation/out/glu-relu-100L_run/checkpoint-100/trainer_state.json ADDED
@@ -0,0 +1,69 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "best_global_step": null,
3
+ "best_metric": null,
4
+ "best_model_checkpoint": null,
5
+ "epoch": 0.006740815638692282,
6
+ "eval_steps": 2498,
7
+ "global_step": 100,
8
+ "is_hyper_param_search": false,
9
+ "is_local_process_zero": true,
10
+ "is_world_process_zero": true,
11
+ "log_history": [
12
+ {
13
+ "epoch": 0.0013481631277384564,
14
+ "grad_norm": 2.09375,
15
+ "learning_rate": 2.66e-05,
16
+ "loss": 8.298905181884766,
17
+ "step": 20
18
+ },
19
+ {
20
+ "epoch": 0.002696326255476913,
21
+ "grad_norm": 1.3359375,
22
+ "learning_rate": 5.46e-05,
23
+ "loss": 8.045733642578124,
24
+ "step": 40
25
+ },
26
+ {
27
+ "epoch": 0.004044489383215369,
28
+ "grad_norm": 1.2734375,
29
+ "learning_rate": 8.259999999999999e-05,
30
+ "loss": 7.839435577392578,
31
+ "step": 60
32
+ },
33
+ {
34
+ "epoch": 0.005392652510953826,
35
+ "grad_norm": 1.265625,
36
+ "learning_rate": 0.0001106,
37
+ "loss": 7.574758148193359,
38
+ "step": 80
39
+ },
40
+ {
41
+ "epoch": 0.006740815638692282,
42
+ "grad_norm": 1.1953125,
43
+ "learning_rate": 0.0001386,
44
+ "loss": 7.235897064208984,
45
+ "step": 100
46
+ }
47
+ ],
48
+ "logging_steps": 20,
49
+ "max_steps": 2500,
50
+ "num_input_tokens_seen": 0,
51
+ "num_train_epochs": 1,
52
+ "save_steps": 100,
53
+ "stateful_callbacks": {
54
+ "TrainerControl": {
55
+ "args": {
56
+ "should_epoch_stop": false,
57
+ "should_evaluate": false,
58
+ "should_log": false,
59
+ "should_save": true,
60
+ "should_training_stop": false
61
+ },
62
+ "attributes": {}
63
+ }
64
+ },
65
+ "total_flos": 322628380262400.0,
66
+ "train_batch_size": 64,
67
+ "trial_name": null,
68
+ "trial_params": null
69
+ }
zain/Activation/out/glu-relu-100L_run/checkpoint-100/training_args.bin ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:13465899109c0aad1422356a6b7684d9c92b0cc10d091b5c3e65ca518136a315
3
+ size 4920
zain/Activation/out/glu-relu-100L_run/checkpoint-1000/config.json ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "activation": "relu",
3
+ "architectures": [
4
+ "TinyLlamaForCausalLM"
5
+ ],
6
+ "attention_bias": false,
7
+ "attention_dropout": 0.0,
8
+ "bos_token_id": 1,
9
+ "dtype": "bfloat16",
10
+ "eos_token_id": 2,
11
+ "head_dim": 32,
12
+ "hidden_act": "silu",
13
+ "hidden_size": 128,
14
+ "initializer_range": 0.02,
15
+ "intermediate_size": 256,
16
+ "max_position_embeddings": 512,
17
+ "mlp_bias": false,
18
+ "mlp_type": "glu",
19
+ "model_type": "tiny_llama",
20
+ "num_attention_heads": 4,
21
+ "num_hidden_layers": 100,
22
+ "num_key_value_heads": 4,
23
+ "pad_token_id": 0,
24
+ "pretraining_tp": 1,
25
+ "rms_norm_eps": 1e-06,
26
+ "rope_parameters": {
27
+ "rope_theta": 10000.0,
28
+ "rope_type": "default"
29
+ },
30
+ "tie_word_embeddings": true,
31
+ "tokenizer_name": "w-ahmad/tiny-stories-tokenizer",
32
+ "transformers_version": "5.16.0.dev0",
33
+ "use_cache": false,
34
+ "vocab_size": 4096,
35
+ "waleed_beta": 4
36
+ }
zain/Activation/out/glu-relu-100L_run/checkpoint-1000/model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:bf3a3764dbb4df21f34d0610c5f2117429ca4a60de6b434cc133ef4a519a0e8c
3
+ size 33967272
zain/Activation/out/glu-relu-100L_run/checkpoint-1000/optimizer.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:51383f388f1162c47afbca62ff7d82aefd08fd24030bda627c9f594104d7373f
3
+ size 68504996
zain/Activation/out/glu-relu-100L_run/checkpoint-1000/rng_state.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:9d9cd6a0487226e5bd30d1846894c82af483733ab4381b75bae9c0745e05d405
3
+ size 14244
zain/Activation/out/glu-relu-100L_run/checkpoint-1000/scheduler.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:cf8d0c79c23f451fbfe25677a50bf40ae55d38cdf93a853c1a033c179996385d
3
+ size 1064
zain/Activation/out/glu-relu-100L_run/checkpoint-1000/tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
zain/Activation/out/glu-relu-100L_run/checkpoint-1000/tokenizer_config.json ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_prefix_space": false,
3
+ "backend": "tokenizers",
4
+ "bos_token": "<|endoftext|>",
5
+ "eos_token": "<|endoftext|>",
6
+ "errors": "replace",
7
+ "is_local": false,
8
+ "local_files_only": false,
9
+ "model_max_length": 1024,
10
+ "pad_token": "<|endoftext|>",
11
+ "tokenizer_class": "GPT2Tokenizer",
12
+ "unk_token": "<|endoftext|>"
13
+ }
zain/Activation/out/glu-relu-100L_run/checkpoint-1000/trainer_state.json ADDED
@@ -0,0 +1,384 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "best_global_step": null,
3
+ "best_metric": null,
4
+ "best_model_checkpoint": null,
5
+ "epoch": 0.06740815638692282,
6
+ "eval_steps": 2498,
7
+ "global_step": 1000,
8
+ "is_hyper_param_search": false,
9
+ "is_local_process_zero": true,
10
+ "is_world_process_zero": true,
11
+ "log_history": [
12
+ {
13
+ "epoch": 0.0013481631277384564,
14
+ "grad_norm": 2.09375,
15
+ "learning_rate": 2.66e-05,
16
+ "loss": 8.298905181884766,
17
+ "step": 20
18
+ },
19
+ {
20
+ "epoch": 0.002696326255476913,
21
+ "grad_norm": 1.3359375,
22
+ "learning_rate": 5.46e-05,
23
+ "loss": 8.045733642578124,
24
+ "step": 40
25
+ },
26
+ {
27
+ "epoch": 0.004044489383215369,
28
+ "grad_norm": 1.2734375,
29
+ "learning_rate": 8.259999999999999e-05,
30
+ "loss": 7.839435577392578,
31
+ "step": 60
32
+ },
33
+ {
34
+ "epoch": 0.005392652510953826,
35
+ "grad_norm": 1.265625,
36
+ "learning_rate": 0.0001106,
37
+ "loss": 7.574758148193359,
38
+ "step": 80
39
+ },
40
+ {
41
+ "epoch": 0.006740815638692282,
42
+ "grad_norm": 1.1953125,
43
+ "learning_rate": 0.0001386,
44
+ "loss": 7.235897064208984,
45
+ "step": 100
46
+ },
47
+ {
48
+ "epoch": 0.008088978766430738,
49
+ "grad_norm": 1.140625,
50
+ "learning_rate": 0.00016659999999999998,
51
+ "loss": 6.862369537353516,
52
+ "step": 120
53
+ },
54
+ {
55
+ "epoch": 0.009437141894169195,
56
+ "grad_norm": 1.0234375,
57
+ "learning_rate": 0.00019460000000000001,
58
+ "loss": 6.508391571044922,
59
+ "step": 140
60
+ },
61
+ {
62
+ "epoch": 0.010785305021907651,
63
+ "grad_norm": 1.078125,
64
+ "learning_rate": 0.0002226,
65
+ "loss": 6.215537261962891,
66
+ "step": 160
67
+ },
68
+ {
69
+ "epoch": 0.012133468149646108,
70
+ "grad_norm": 0.76171875,
71
+ "learning_rate": 0.00025059999999999997,
72
+ "loss": 5.984259414672851,
73
+ "step": 180
74
+ },
75
+ {
76
+ "epoch": 0.013481631277384564,
77
+ "grad_norm": 0.82421875,
78
+ "learning_rate": 0.0002786,
79
+ "loss": 5.756989669799805,
80
+ "step": 200
81
+ },
82
+ {
83
+ "epoch": 0.01482979440512302,
84
+ "grad_norm": 1.171875,
85
+ "learning_rate": 0.00030659999999999997,
86
+ "loss": 5.531437683105469,
87
+ "step": 220
88
+ },
89
+ {
90
+ "epoch": 0.016177957532861477,
91
+ "grad_norm": 1.1171875,
92
+ "learning_rate": 0.0003346,
93
+ "loss": 5.294247817993164,
94
+ "step": 240
95
+ },
96
+ {
97
+ "epoch": 0.01752612066059993,
98
+ "grad_norm": 0.9453125,
99
+ "learning_rate": 0.00036260000000000003,
100
+ "loss": 5.055661392211914,
101
+ "step": 260
102
+ },
103
+ {
104
+ "epoch": 0.01887428378833839,
105
+ "grad_norm": 0.94921875,
106
+ "learning_rate": 0.0003906,
107
+ "loss": 4.797793197631836,
108
+ "step": 280
109
+ },
110
+ {
111
+ "epoch": 0.020222446916076844,
112
+ "grad_norm": 1.421875,
113
+ "learning_rate": 0.0004186,
114
+ "loss": 4.594622802734375,
115
+ "step": 300
116
+ },
117
+ {
118
+ "epoch": 0.021570610043815303,
119
+ "grad_norm": 1.2265625,
120
+ "learning_rate": 0.0004466,
121
+ "loss": 4.39941635131836,
122
+ "step": 320
123
+ },
124
+ {
125
+ "epoch": 0.022918773171553757,
126
+ "grad_norm": 0.5546875,
127
+ "learning_rate": 0.00047460000000000004,
128
+ "loss": 4.24058723449707,
129
+ "step": 340
130
+ },
131
+ {
132
+ "epoch": 0.024266936299292215,
133
+ "grad_norm": 0.83984375,
134
+ "learning_rate": 0.0005026,
135
+ "loss": 4.102418899536133,
136
+ "step": 360
137
+ },
138
+ {
139
+ "epoch": 0.02561509942703067,
140
+ "grad_norm": 0.70703125,
141
+ "learning_rate": 0.0005306,
142
+ "loss": 3.975328063964844,
143
+ "step": 380
144
+ },
145
+ {
146
+ "epoch": 0.026963262554769128,
147
+ "grad_norm": 0.6171875,
148
+ "learning_rate": 0.0005586,
149
+ "loss": 3.839755630493164,
150
+ "step": 400
151
+ },
152
+ {
153
+ "epoch": 0.028311425682507583,
154
+ "grad_norm": 0.91796875,
155
+ "learning_rate": 0.0005866,
156
+ "loss": 3.7486804962158202,
157
+ "step": 420
158
+ },
159
+ {
160
+ "epoch": 0.02965958881024604,
161
+ "grad_norm": 0.61328125,
162
+ "learning_rate": 0.0006146,
163
+ "loss": 3.6831081390380858,
164
+ "step": 440
165
+ },
166
+ {
167
+ "epoch": 0.031007751937984496,
168
+ "grad_norm": 0.875,
169
+ "learning_rate": 0.0006426,
170
+ "loss": 3.5874557495117188,
171
+ "step": 460
172
+ },
173
+ {
174
+ "epoch": 0.032355915065722954,
175
+ "grad_norm": 0.546875,
176
+ "learning_rate": 0.0006705999999999999,
177
+ "loss": 3.5174354553222655,
178
+ "step": 480
179
+ },
180
+ {
181
+ "epoch": 0.03370407819346141,
182
+ "grad_norm": 0.66796875,
183
+ "learning_rate": 0.0006986,
184
+ "loss": 3.439300537109375,
185
+ "step": 500
186
+ },
187
+ {
188
+ "epoch": 0.03505224132119986,
189
+ "grad_norm": 0.54296875,
190
+ "learning_rate": 0.0007,
191
+ "loss": 3.386067581176758,
192
+ "step": 520
193
+ },
194
+ {
195
+ "epoch": 0.03640040444893832,
196
+ "grad_norm": 0.6015625,
197
+ "learning_rate": 0.0007,
198
+ "loss": 3.3094718933105467,
199
+ "step": 540
200
+ },
201
+ {
202
+ "epoch": 0.03774856757667678,
203
+ "grad_norm": 0.55859375,
204
+ "learning_rate": 0.0007,
205
+ "loss": 3.2609573364257813,
206
+ "step": 560
207
+ },
208
+ {
209
+ "epoch": 0.03909673070441524,
210
+ "grad_norm": 0.61328125,
211
+ "learning_rate": 0.0007,
212
+ "loss": 3.2046344757080076,
213
+ "step": 580
214
+ },
215
+ {
216
+ "epoch": 0.04044489383215369,
217
+ "grad_norm": 0.5234375,
218
+ "learning_rate": 0.0007,
219
+ "loss": 3.17300968170166,
220
+ "step": 600
221
+ },
222
+ {
223
+ "epoch": 0.04179305695989215,
224
+ "grad_norm": 0.578125,
225
+ "learning_rate": 0.0007,
226
+ "loss": 3.1147342681884767,
227
+ "step": 620
228
+ },
229
+ {
230
+ "epoch": 0.043141220087630605,
231
+ "grad_norm": 0.4765625,
232
+ "learning_rate": 0.0007,
233
+ "loss": 3.0973543167114257,
234
+ "step": 640
235
+ },
236
+ {
237
+ "epoch": 0.044489383215369056,
238
+ "grad_norm": 0.453125,
239
+ "learning_rate": 0.0007,
240
+ "loss": 3.043070602416992,
241
+ "step": 660
242
+ },
243
+ {
244
+ "epoch": 0.045837546343107514,
245
+ "grad_norm": 0.455078125,
246
+ "learning_rate": 0.0007,
247
+ "loss": 3.0213077545166014,
248
+ "step": 680
249
+ },
250
+ {
251
+ "epoch": 0.04718570947084597,
252
+ "grad_norm": 0.419921875,
253
+ "learning_rate": 0.0007,
254
+ "loss": 2.977329063415527,
255
+ "step": 700
256
+ },
257
+ {
258
+ "epoch": 0.04853387259858443,
259
+ "grad_norm": 0.490234375,
260
+ "learning_rate": 0.0007,
261
+ "loss": 2.942974090576172,
262
+ "step": 720
263
+ },
264
+ {
265
+ "epoch": 0.04988203572632288,
266
+ "grad_norm": 0.48046875,
267
+ "learning_rate": 0.0007,
268
+ "loss": 2.9064186096191404,
269
+ "step": 740
270
+ },
271
+ {
272
+ "epoch": 0.05123019885406134,
273
+ "grad_norm": 0.498046875,
274
+ "learning_rate": 0.0007,
275
+ "loss": 2.890872001647949,
276
+ "step": 760
277
+ },
278
+ {
279
+ "epoch": 0.0525783619817998,
280
+ "grad_norm": 0.5,
281
+ "learning_rate": 0.0007,
282
+ "loss": 2.8978702545166017,
283
+ "step": 780
284
+ },
285
+ {
286
+ "epoch": 0.053926525109538256,
287
+ "grad_norm": 0.466796875,
288
+ "learning_rate": 0.0007,
289
+ "loss": 2.8463741302490235,
290
+ "step": 800
291
+ },
292
+ {
293
+ "epoch": 0.05527468823727671,
294
+ "grad_norm": 0.47265625,
295
+ "learning_rate": 0.0007,
296
+ "loss": 2.8323719024658205,
297
+ "step": 820
298
+ },
299
+ {
300
+ "epoch": 0.056622851365015166,
301
+ "grad_norm": 0.5,
302
+ "learning_rate": 0.0007,
303
+ "loss": 2.8007469177246094,
304
+ "step": 840
305
+ },
306
+ {
307
+ "epoch": 0.057971014492753624,
308
+ "grad_norm": 0.43359375,
309
+ "learning_rate": 0.0007,
310
+ "loss": 2.7573240280151365,
311
+ "step": 860
312
+ },
313
+ {
314
+ "epoch": 0.05931917762049208,
315
+ "grad_norm": 0.470703125,
316
+ "learning_rate": 0.0007,
317
+ "loss": 2.7341793060302733,
318
+ "step": 880
319
+ },
320
+ {
321
+ "epoch": 0.06066734074823053,
322
+ "grad_norm": 0.546875,
323
+ "learning_rate": 0.0007,
324
+ "loss": 2.7180759429931642,
325
+ "step": 900
326
+ },
327
+ {
328
+ "epoch": 0.06201550387596899,
329
+ "grad_norm": 0.5625,
330
+ "learning_rate": 0.0007,
331
+ "loss": 2.692394256591797,
332
+ "step": 920
333
+ },
334
+ {
335
+ "epoch": 0.06336366700370745,
336
+ "grad_norm": 0.443359375,
337
+ "learning_rate": 0.0007,
338
+ "loss": 2.678949546813965,
339
+ "step": 940
340
+ },
341
+ {
342
+ "epoch": 0.06471183013144591,
343
+ "grad_norm": 0.466796875,
344
+ "learning_rate": 0.0007,
345
+ "loss": 2.651218223571777,
346
+ "step": 960
347
+ },
348
+ {
349
+ "epoch": 0.06605999325918437,
350
+ "grad_norm": 0.43359375,
351
+ "learning_rate": 0.0007,
352
+ "loss": 2.632099723815918,
353
+ "step": 980
354
+ },
355
+ {
356
+ "epoch": 0.06740815638692282,
357
+ "grad_norm": 0.470703125,
358
+ "learning_rate": 0.0007,
359
+ "loss": 2.6305675506591797,
360
+ "step": 1000
361
+ }
362
+ ],
363
+ "logging_steps": 20,
364
+ "max_steps": 2500,
365
+ "num_input_tokens_seen": 0,
366
+ "num_train_epochs": 1,
367
+ "save_steps": 100,
368
+ "stateful_callbacks": {
369
+ "TrainerControl": {
370
+ "args": {
371
+ "should_epoch_stop": false,
372
+ "should_evaluate": false,
373
+ "should_log": false,
374
+ "should_save": true,
375
+ "should_training_stop": false
376
+ },
377
+ "attributes": {}
378
+ }
379
+ },
380
+ "total_flos": 3226283802624000.0,
381
+ "train_batch_size": 64,
382
+ "trial_name": null,
383
+ "trial_params": null
384
+ }
zain/Activation/out/glu-relu-100L_run/checkpoint-1000/training_args.bin ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:13465899109c0aad1422356a6b7684d9c92b0cc10d091b5c3e65ca518136a315
3
+ size 4920
zain/Activation/out/glu-relu-100L_run/checkpoint-1100/config.json ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "activation": "relu",
3
+ "architectures": [
4
+ "TinyLlamaForCausalLM"
5
+ ],
6
+ "attention_bias": false,
7
+ "attention_dropout": 0.0,
8
+ "bos_token_id": 1,
9
+ "dtype": "bfloat16",
10
+ "eos_token_id": 2,
11
+ "head_dim": 32,
12
+ "hidden_act": "silu",
13
+ "hidden_size": 128,
14
+ "initializer_range": 0.02,
15
+ "intermediate_size": 256,
16
+ "max_position_embeddings": 512,
17
+ "mlp_bias": false,
18
+ "mlp_type": "glu",
19
+ "model_type": "tiny_llama",
20
+ "num_attention_heads": 4,
21
+ "num_hidden_layers": 100,
22
+ "num_key_value_heads": 4,
23
+ "pad_token_id": 0,
24
+ "pretraining_tp": 1,
25
+ "rms_norm_eps": 1e-06,
26
+ "rope_parameters": {
27
+ "rope_theta": 10000.0,
28
+ "rope_type": "default"
29
+ },
30
+ "tie_word_embeddings": true,
31
+ "tokenizer_name": "w-ahmad/tiny-stories-tokenizer",
32
+ "transformers_version": "5.16.0.dev0",
33
+ "use_cache": false,
34
+ "vocab_size": 4096,
35
+ "waleed_beta": 4
36
+ }
zain/Activation/out/glu-relu-100L_run/checkpoint-1100/model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:f1521d1da189673ea1656d53a9ee6fefad9281dc14b2b385983acebedd913e4d
3
+ size 33967272
zain/Activation/out/glu-relu-100L_run/checkpoint-1100/optimizer.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d46d83c2365e0d5242337e6b2909e950ccec97226b06fc627dfc24372366d900
3
+ size 68504996
zain/Activation/out/glu-relu-100L_run/checkpoint-1100/rng_state.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:9d9cd6a0487226e5bd30d1846894c82af483733ab4381b75bae9c0745e05d405
3
+ size 14244
zain/Activation/out/glu-relu-100L_run/checkpoint-1100/scheduler.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:88121baf7de31688d97eca5c32cb058d585673d52f936d32dcd04e8b1903dcac
3
+ size 1064
zain/Activation/out/glu-relu-100L_run/checkpoint-1100/tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
zain/Activation/out/glu-relu-100L_run/checkpoint-1100/tokenizer_config.json ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_prefix_space": false,
3
+ "backend": "tokenizers",
4
+ "bos_token": "<|endoftext|>",
5
+ "eos_token": "<|endoftext|>",
6
+ "errors": "replace",
7
+ "is_local": false,
8
+ "local_files_only": false,
9
+ "model_max_length": 1024,
10
+ "pad_token": "<|endoftext|>",
11
+ "tokenizer_class": "GPT2Tokenizer",
12
+ "unk_token": "<|endoftext|>"
13
+ }
zain/Activation/out/glu-relu-100L_run/checkpoint-1100/trainer_state.json ADDED
@@ -0,0 +1,419 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "best_global_step": null,
3
+ "best_metric": null,
4
+ "best_model_checkpoint": null,
5
+ "epoch": 0.0741489720256151,
6
+ "eval_steps": 2498,
7
+ "global_step": 1100,
8
+ "is_hyper_param_search": false,
9
+ "is_local_process_zero": true,
10
+ "is_world_process_zero": true,
11
+ "log_history": [
12
+ {
13
+ "epoch": 0.0013481631277384564,
14
+ "grad_norm": 2.09375,
15
+ "learning_rate": 2.66e-05,
16
+ "loss": 8.298905181884766,
17
+ "step": 20
18
+ },
19
+ {
20
+ "epoch": 0.002696326255476913,
21
+ "grad_norm": 1.3359375,
22
+ "learning_rate": 5.46e-05,
23
+ "loss": 8.045733642578124,
24
+ "step": 40
25
+ },
26
+ {
27
+ "epoch": 0.004044489383215369,
28
+ "grad_norm": 1.2734375,
29
+ "learning_rate": 8.259999999999999e-05,
30
+ "loss": 7.839435577392578,
31
+ "step": 60
32
+ },
33
+ {
34
+ "epoch": 0.005392652510953826,
35
+ "grad_norm": 1.265625,
36
+ "learning_rate": 0.0001106,
37
+ "loss": 7.574758148193359,
38
+ "step": 80
39
+ },
40
+ {
41
+ "epoch": 0.006740815638692282,
42
+ "grad_norm": 1.1953125,
43
+ "learning_rate": 0.0001386,
44
+ "loss": 7.235897064208984,
45
+ "step": 100
46
+ },
47
+ {
48
+ "epoch": 0.008088978766430738,
49
+ "grad_norm": 1.140625,
50
+ "learning_rate": 0.00016659999999999998,
51
+ "loss": 6.862369537353516,
52
+ "step": 120
53
+ },
54
+ {
55
+ "epoch": 0.009437141894169195,
56
+ "grad_norm": 1.0234375,
57
+ "learning_rate": 0.00019460000000000001,
58
+ "loss": 6.508391571044922,
59
+ "step": 140
60
+ },
61
+ {
62
+ "epoch": 0.010785305021907651,
63
+ "grad_norm": 1.078125,
64
+ "learning_rate": 0.0002226,
65
+ "loss": 6.215537261962891,
66
+ "step": 160
67
+ },
68
+ {
69
+ "epoch": 0.012133468149646108,
70
+ "grad_norm": 0.76171875,
71
+ "learning_rate": 0.00025059999999999997,
72
+ "loss": 5.984259414672851,
73
+ "step": 180
74
+ },
75
+ {
76
+ "epoch": 0.013481631277384564,
77
+ "grad_norm": 0.82421875,
78
+ "learning_rate": 0.0002786,
79
+ "loss": 5.756989669799805,
80
+ "step": 200
81
+ },
82
+ {
83
+ "epoch": 0.01482979440512302,
84
+ "grad_norm": 1.171875,
85
+ "learning_rate": 0.00030659999999999997,
86
+ "loss": 5.531437683105469,
87
+ "step": 220
88
+ },
89
+ {
90
+ "epoch": 0.016177957532861477,
91
+ "grad_norm": 1.1171875,
92
+ "learning_rate": 0.0003346,
93
+ "loss": 5.294247817993164,
94
+ "step": 240
95
+ },
96
+ {
97
+ "epoch": 0.01752612066059993,
98
+ "grad_norm": 0.9453125,
99
+ "learning_rate": 0.00036260000000000003,
100
+ "loss": 5.055661392211914,
101
+ "step": 260
102
+ },
103
+ {
104
+ "epoch": 0.01887428378833839,
105
+ "grad_norm": 0.94921875,
106
+ "learning_rate": 0.0003906,
107
+ "loss": 4.797793197631836,
108
+ "step": 280
109
+ },
110
+ {
111
+ "epoch": 0.020222446916076844,
112
+ "grad_norm": 1.421875,
113
+ "learning_rate": 0.0004186,
114
+ "loss": 4.594622802734375,
115
+ "step": 300
116
+ },
117
+ {
118
+ "epoch": 0.021570610043815303,
119
+ "grad_norm": 1.2265625,
120
+ "learning_rate": 0.0004466,
121
+ "loss": 4.39941635131836,
122
+ "step": 320
123
+ },
124
+ {
125
+ "epoch": 0.022918773171553757,
126
+ "grad_norm": 0.5546875,
127
+ "learning_rate": 0.00047460000000000004,
128
+ "loss": 4.24058723449707,
129
+ "step": 340
130
+ },
131
+ {
132
+ "epoch": 0.024266936299292215,
133
+ "grad_norm": 0.83984375,
134
+ "learning_rate": 0.0005026,
135
+ "loss": 4.102418899536133,
136
+ "step": 360
137
+ },
138
+ {
139
+ "epoch": 0.02561509942703067,
140
+ "grad_norm": 0.70703125,
141
+ "learning_rate": 0.0005306,
142
+ "loss": 3.975328063964844,
143
+ "step": 380
144
+ },
145
+ {
146
+ "epoch": 0.026963262554769128,
147
+ "grad_norm": 0.6171875,
148
+ "learning_rate": 0.0005586,
149
+ "loss": 3.839755630493164,
150
+ "step": 400
151
+ },
152
+ {
153
+ "epoch": 0.028311425682507583,
154
+ "grad_norm": 0.91796875,
155
+ "learning_rate": 0.0005866,
156
+ "loss": 3.7486804962158202,
157
+ "step": 420
158
+ },
159
+ {
160
+ "epoch": 0.02965958881024604,
161
+ "grad_norm": 0.61328125,
162
+ "learning_rate": 0.0006146,
163
+ "loss": 3.6831081390380858,
164
+ "step": 440
165
+ },
166
+ {
167
+ "epoch": 0.031007751937984496,
168
+ "grad_norm": 0.875,
169
+ "learning_rate": 0.0006426,
170
+ "loss": 3.5874557495117188,
171
+ "step": 460
172
+ },
173
+ {
174
+ "epoch": 0.032355915065722954,
175
+ "grad_norm": 0.546875,
176
+ "learning_rate": 0.0006705999999999999,
177
+ "loss": 3.5174354553222655,
178
+ "step": 480
179
+ },
180
+ {
181
+ "epoch": 0.03370407819346141,
182
+ "grad_norm": 0.66796875,
183
+ "learning_rate": 0.0006986,
184
+ "loss": 3.439300537109375,
185
+ "step": 500
186
+ },
187
+ {
188
+ "epoch": 0.03505224132119986,
189
+ "grad_norm": 0.54296875,
190
+ "learning_rate": 0.0007,
191
+ "loss": 3.386067581176758,
192
+ "step": 520
193
+ },
194
+ {
195
+ "epoch": 0.03640040444893832,
196
+ "grad_norm": 0.6015625,
197
+ "learning_rate": 0.0007,
198
+ "loss": 3.3094718933105467,
199
+ "step": 540
200
+ },
201
+ {
202
+ "epoch": 0.03774856757667678,
203
+ "grad_norm": 0.55859375,
204
+ "learning_rate": 0.0007,
205
+ "loss": 3.2609573364257813,
206
+ "step": 560
207
+ },
208
+ {
209
+ "epoch": 0.03909673070441524,
210
+ "grad_norm": 0.61328125,
211
+ "learning_rate": 0.0007,
212
+ "loss": 3.2046344757080076,
213
+ "step": 580
214
+ },
215
+ {
216
+ "epoch": 0.04044489383215369,
217
+ "grad_norm": 0.5234375,
218
+ "learning_rate": 0.0007,
219
+ "loss": 3.17300968170166,
220
+ "step": 600
221
+ },
222
+ {
223
+ "epoch": 0.04179305695989215,
224
+ "grad_norm": 0.578125,
225
+ "learning_rate": 0.0007,
226
+ "loss": 3.1147342681884767,
227
+ "step": 620
228
+ },
229
+ {
230
+ "epoch": 0.043141220087630605,
231
+ "grad_norm": 0.4765625,
232
+ "learning_rate": 0.0007,
233
+ "loss": 3.0973543167114257,
234
+ "step": 640
235
+ },
236
+ {
237
+ "epoch": 0.044489383215369056,
238
+ "grad_norm": 0.453125,
239
+ "learning_rate": 0.0007,
240
+ "loss": 3.043070602416992,
241
+ "step": 660
242
+ },
243
+ {
244
+ "epoch": 0.045837546343107514,
245
+ "grad_norm": 0.455078125,
246
+ "learning_rate": 0.0007,
247
+ "loss": 3.0213077545166014,
248
+ "step": 680
249
+ },
250
+ {
251
+ "epoch": 0.04718570947084597,
252
+ "grad_norm": 0.419921875,
253
+ "learning_rate": 0.0007,
254
+ "loss": 2.977329063415527,
255
+ "step": 700
256
+ },
257
+ {
258
+ "epoch": 0.04853387259858443,
259
+ "grad_norm": 0.490234375,
260
+ "learning_rate": 0.0007,
261
+ "loss": 2.942974090576172,
262
+ "step": 720
263
+ },
264
+ {
265
+ "epoch": 0.04988203572632288,
266
+ "grad_norm": 0.48046875,
267
+ "learning_rate": 0.0007,
268
+ "loss": 2.9064186096191404,
269
+ "step": 740
270
+ },
271
+ {
272
+ "epoch": 0.05123019885406134,
273
+ "grad_norm": 0.498046875,
274
+ "learning_rate": 0.0007,
275
+ "loss": 2.890872001647949,
276
+ "step": 760
277
+ },
278
+ {
279
+ "epoch": 0.0525783619817998,
280
+ "grad_norm": 0.5,
281
+ "learning_rate": 0.0007,
282
+ "loss": 2.8978702545166017,
283
+ "step": 780
284
+ },
285
+ {
286
+ "epoch": 0.053926525109538256,
287
+ "grad_norm": 0.466796875,
288
+ "learning_rate": 0.0007,
289
+ "loss": 2.8463741302490235,
290
+ "step": 800
291
+ },
292
+ {
293
+ "epoch": 0.05527468823727671,
294
+ "grad_norm": 0.47265625,
295
+ "learning_rate": 0.0007,
296
+ "loss": 2.8323719024658205,
297
+ "step": 820
298
+ },
299
+ {
300
+ "epoch": 0.056622851365015166,
301
+ "grad_norm": 0.5,
302
+ "learning_rate": 0.0007,
303
+ "loss": 2.8007469177246094,
304
+ "step": 840
305
+ },
306
+ {
307
+ "epoch": 0.057971014492753624,
308
+ "grad_norm": 0.43359375,
309
+ "learning_rate": 0.0007,
310
+ "loss": 2.7573240280151365,
311
+ "step": 860
312
+ },
313
+ {
314
+ "epoch": 0.05931917762049208,
315
+ "grad_norm": 0.470703125,
316
+ "learning_rate": 0.0007,
317
+ "loss": 2.7341793060302733,
318
+ "step": 880
319
+ },
320
+ {
321
+ "epoch": 0.06066734074823053,
322
+ "grad_norm": 0.546875,
323
+ "learning_rate": 0.0007,
324
+ "loss": 2.7180759429931642,
325
+ "step": 900
326
+ },
327
+ {
328
+ "epoch": 0.06201550387596899,
329
+ "grad_norm": 0.5625,
330
+ "learning_rate": 0.0007,
331
+ "loss": 2.692394256591797,
332
+ "step": 920
333
+ },
334
+ {
335
+ "epoch": 0.06336366700370745,
336
+ "grad_norm": 0.443359375,
337
+ "learning_rate": 0.0007,
338
+ "loss": 2.678949546813965,
339
+ "step": 940
340
+ },
341
+ {
342
+ "epoch": 0.06471183013144591,
343
+ "grad_norm": 0.466796875,
344
+ "learning_rate": 0.0007,
345
+ "loss": 2.651218223571777,
346
+ "step": 960
347
+ },
348
+ {
349
+ "epoch": 0.06605999325918437,
350
+ "grad_norm": 0.43359375,
351
+ "learning_rate": 0.0007,
352
+ "loss": 2.632099723815918,
353
+ "step": 980
354
+ },
355
+ {
356
+ "epoch": 0.06740815638692282,
357
+ "grad_norm": 0.470703125,
358
+ "learning_rate": 0.0007,
359
+ "loss": 2.6305675506591797,
360
+ "step": 1000
361
+ },
362
+ {
363
+ "epoch": 0.06875631951466127,
364
+ "grad_norm": 0.462890625,
365
+ "learning_rate": 0.0007,
366
+ "loss": 2.5994905471801757,
367
+ "step": 1020
368
+ },
369
+ {
370
+ "epoch": 0.07010448264239973,
371
+ "grad_norm": 0.474609375,
372
+ "learning_rate": 0.0007,
373
+ "loss": 2.5834341049194336,
374
+ "step": 1040
375
+ },
376
+ {
377
+ "epoch": 0.07145264577013818,
378
+ "grad_norm": 0.498046875,
379
+ "learning_rate": 0.0007,
380
+ "loss": 2.567459487915039,
381
+ "step": 1060
382
+ },
383
+ {
384
+ "epoch": 0.07280080889787664,
385
+ "grad_norm": 0.462890625,
386
+ "learning_rate": 0.0007,
387
+ "loss": 2.54827823638916,
388
+ "step": 1080
389
+ },
390
+ {
391
+ "epoch": 0.0741489720256151,
392
+ "grad_norm": 0.478515625,
393
+ "learning_rate": 0.0007,
394
+ "loss": 2.5431196212768556,
395
+ "step": 1100
396
+ }
397
+ ],
398
+ "logging_steps": 20,
399
+ "max_steps": 2500,
400
+ "num_input_tokens_seen": 0,
401
+ "num_train_epochs": 1,
402
+ "save_steps": 100,
403
+ "stateful_callbacks": {
404
+ "TrainerControl": {
405
+ "args": {
406
+ "should_epoch_stop": false,
407
+ "should_evaluate": false,
408
+ "should_log": false,
409
+ "should_save": true,
410
+ "should_training_stop": false
411
+ },
412
+ "attributes": {}
413
+ }
414
+ },
415
+ "total_flos": 3548912182886400.0,
416
+ "train_batch_size": 64,
417
+ "trial_name": null,
418
+ "trial_params": null
419
+ }
zain/Activation/out/glu-relu-100L_run/checkpoint-1100/training_args.bin ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:13465899109c0aad1422356a6b7684d9c92b0cc10d091b5c3e65ca518136a315
3
+ size 4920
zain/Activation/out/glu-relu-100L_run/checkpoint-1200/config.json ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "activation": "relu",
3
+ "architectures": [
4
+ "TinyLlamaForCausalLM"
5
+ ],
6
+ "attention_bias": false,
7
+ "attention_dropout": 0.0,
8
+ "bos_token_id": 1,
9
+ "dtype": "bfloat16",
10
+ "eos_token_id": 2,
11
+ "head_dim": 32,
12
+ "hidden_act": "silu",
13
+ "hidden_size": 128,
14
+ "initializer_range": 0.02,
15
+ "intermediate_size": 256,
16
+ "max_position_embeddings": 512,
17
+ "mlp_bias": false,
18
+ "mlp_type": "glu",
19
+ "model_type": "tiny_llama",
20
+ "num_attention_heads": 4,
21
+ "num_hidden_layers": 100,
22
+ "num_key_value_heads": 4,
23
+ "pad_token_id": 0,
24
+ "pretraining_tp": 1,
25
+ "rms_norm_eps": 1e-06,
26
+ "rope_parameters": {
27
+ "rope_theta": 10000.0,
28
+ "rope_type": "default"
29
+ },
30
+ "tie_word_embeddings": true,
31
+ "tokenizer_name": "w-ahmad/tiny-stories-tokenizer",
32
+ "transformers_version": "5.16.0.dev0",
33
+ "use_cache": false,
34
+ "vocab_size": 4096,
35
+ "waleed_beta": 4
36
+ }
zain/Activation/out/glu-relu-100L_run/checkpoint-1200/model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:8cc67161d097e3c2c21fe51c93cc7efae656d4746264785734b48960aed1c4cf
3
+ size 33967272
zain/Activation/out/glu-relu-100L_run/checkpoint-1200/optimizer.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:2fb3df0166cb4535efb45fcc363fc6402a29359a750439360b9a03e753639d36
3
+ size 68504996
zain/Activation/out/glu-relu-100L_run/checkpoint-1200/rng_state.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:9d9cd6a0487226e5bd30d1846894c82af483733ab4381b75bae9c0745e05d405
3
+ size 14244
zain/Activation/out/glu-relu-100L_run/checkpoint-1200/scheduler.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:8543afc9166865219cadbba29db865d839ec4102a4f9ddeffc962b9fdba3b028
3
+ size 1064
zain/Activation/out/glu-relu-100L_run/checkpoint-1200/tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
zain/Activation/out/glu-relu-100L_run/checkpoint-1200/tokenizer_config.json ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_prefix_space": false,
3
+ "backend": "tokenizers",
4
+ "bos_token": "<|endoftext|>",
5
+ "eos_token": "<|endoftext|>",
6
+ "errors": "replace",
7
+ "is_local": false,
8
+ "local_files_only": false,
9
+ "model_max_length": 1024,
10
+ "pad_token": "<|endoftext|>",
11
+ "tokenizer_class": "GPT2Tokenizer",
12
+ "unk_token": "<|endoftext|>"
13
+ }
zain/Activation/out/glu-relu-100L_run/checkpoint-1200/trainer_state.json ADDED
@@ -0,0 +1,454 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "best_global_step": null,
3
+ "best_metric": null,
4
+ "best_model_checkpoint": null,
5
+ "epoch": 0.08088978766430738,
6
+ "eval_steps": 2498,
7
+ "global_step": 1200,
8
+ "is_hyper_param_search": false,
9
+ "is_local_process_zero": true,
10
+ "is_world_process_zero": true,
11
+ "log_history": [
12
+ {
13
+ "epoch": 0.0013481631277384564,
14
+ "grad_norm": 2.09375,
15
+ "learning_rate": 2.66e-05,
16
+ "loss": 8.298905181884766,
17
+ "step": 20
18
+ },
19
+ {
20
+ "epoch": 0.002696326255476913,
21
+ "grad_norm": 1.3359375,
22
+ "learning_rate": 5.46e-05,
23
+ "loss": 8.045733642578124,
24
+ "step": 40
25
+ },
26
+ {
27
+ "epoch": 0.004044489383215369,
28
+ "grad_norm": 1.2734375,
29
+ "learning_rate": 8.259999999999999e-05,
30
+ "loss": 7.839435577392578,
31
+ "step": 60
32
+ },
33
+ {
34
+ "epoch": 0.005392652510953826,
35
+ "grad_norm": 1.265625,
36
+ "learning_rate": 0.0001106,
37
+ "loss": 7.574758148193359,
38
+ "step": 80
39
+ },
40
+ {
41
+ "epoch": 0.006740815638692282,
42
+ "grad_norm": 1.1953125,
43
+ "learning_rate": 0.0001386,
44
+ "loss": 7.235897064208984,
45
+ "step": 100
46
+ },
47
+ {
48
+ "epoch": 0.008088978766430738,
49
+ "grad_norm": 1.140625,
50
+ "learning_rate": 0.00016659999999999998,
51
+ "loss": 6.862369537353516,
52
+ "step": 120
53
+ },
54
+ {
55
+ "epoch": 0.009437141894169195,
56
+ "grad_norm": 1.0234375,
57
+ "learning_rate": 0.00019460000000000001,
58
+ "loss": 6.508391571044922,
59
+ "step": 140
60
+ },
61
+ {
62
+ "epoch": 0.010785305021907651,
63
+ "grad_norm": 1.078125,
64
+ "learning_rate": 0.0002226,
65
+ "loss": 6.215537261962891,
66
+ "step": 160
67
+ },
68
+ {
69
+ "epoch": 0.012133468149646108,
70
+ "grad_norm": 0.76171875,
71
+ "learning_rate": 0.00025059999999999997,
72
+ "loss": 5.984259414672851,
73
+ "step": 180
74
+ },
75
+ {
76
+ "epoch": 0.013481631277384564,
77
+ "grad_norm": 0.82421875,
78
+ "learning_rate": 0.0002786,
79
+ "loss": 5.756989669799805,
80
+ "step": 200
81
+ },
82
+ {
83
+ "epoch": 0.01482979440512302,
84
+ "grad_norm": 1.171875,
85
+ "learning_rate": 0.00030659999999999997,
86
+ "loss": 5.531437683105469,
87
+ "step": 220
88
+ },
89
+ {
90
+ "epoch": 0.016177957532861477,
91
+ "grad_norm": 1.1171875,
92
+ "learning_rate": 0.0003346,
93
+ "loss": 5.294247817993164,
94
+ "step": 240
95
+ },
96
+ {
97
+ "epoch": 0.01752612066059993,
98
+ "grad_norm": 0.9453125,
99
+ "learning_rate": 0.00036260000000000003,
100
+ "loss": 5.055661392211914,
101
+ "step": 260
102
+ },
103
+ {
104
+ "epoch": 0.01887428378833839,
105
+ "grad_norm": 0.94921875,
106
+ "learning_rate": 0.0003906,
107
+ "loss": 4.797793197631836,
108
+ "step": 280
109
+ },
110
+ {
111
+ "epoch": 0.020222446916076844,
112
+ "grad_norm": 1.421875,
113
+ "learning_rate": 0.0004186,
114
+ "loss": 4.594622802734375,
115
+ "step": 300
116
+ },
117
+ {
118
+ "epoch": 0.021570610043815303,
119
+ "grad_norm": 1.2265625,
120
+ "learning_rate": 0.0004466,
121
+ "loss": 4.39941635131836,
122
+ "step": 320
123
+ },
124
+ {
125
+ "epoch": 0.022918773171553757,
126
+ "grad_norm": 0.5546875,
127
+ "learning_rate": 0.00047460000000000004,
128
+ "loss": 4.24058723449707,
129
+ "step": 340
130
+ },
131
+ {
132
+ "epoch": 0.024266936299292215,
133
+ "grad_norm": 0.83984375,
134
+ "learning_rate": 0.0005026,
135
+ "loss": 4.102418899536133,
136
+ "step": 360
137
+ },
138
+ {
139
+ "epoch": 0.02561509942703067,
140
+ "grad_norm": 0.70703125,
141
+ "learning_rate": 0.0005306,
142
+ "loss": 3.975328063964844,
143
+ "step": 380
144
+ },
145
+ {
146
+ "epoch": 0.026963262554769128,
147
+ "grad_norm": 0.6171875,
148
+ "learning_rate": 0.0005586,
149
+ "loss": 3.839755630493164,
150
+ "step": 400
151
+ },
152
+ {
153
+ "epoch": 0.028311425682507583,
154
+ "grad_norm": 0.91796875,
155
+ "learning_rate": 0.0005866,
156
+ "loss": 3.7486804962158202,
157
+ "step": 420
158
+ },
159
+ {
160
+ "epoch": 0.02965958881024604,
161
+ "grad_norm": 0.61328125,
162
+ "learning_rate": 0.0006146,
163
+ "loss": 3.6831081390380858,
164
+ "step": 440
165
+ },
166
+ {
167
+ "epoch": 0.031007751937984496,
168
+ "grad_norm": 0.875,
169
+ "learning_rate": 0.0006426,
170
+ "loss": 3.5874557495117188,
171
+ "step": 460
172
+ },
173
+ {
174
+ "epoch": 0.032355915065722954,
175
+ "grad_norm": 0.546875,
176
+ "learning_rate": 0.0006705999999999999,
177
+ "loss": 3.5174354553222655,
178
+ "step": 480
179
+ },
180
+ {
181
+ "epoch": 0.03370407819346141,
182
+ "grad_norm": 0.66796875,
183
+ "learning_rate": 0.0006986,
184
+ "loss": 3.439300537109375,
185
+ "step": 500
186
+ },
187
+ {
188
+ "epoch": 0.03505224132119986,
189
+ "grad_norm": 0.54296875,
190
+ "learning_rate": 0.0007,
191
+ "loss": 3.386067581176758,
192
+ "step": 520
193
+ },
194
+ {
195
+ "epoch": 0.03640040444893832,
196
+ "grad_norm": 0.6015625,
197
+ "learning_rate": 0.0007,
198
+ "loss": 3.3094718933105467,
199
+ "step": 540
200
+ },
201
+ {
202
+ "epoch": 0.03774856757667678,
203
+ "grad_norm": 0.55859375,
204
+ "learning_rate": 0.0007,
205
+ "loss": 3.2609573364257813,
206
+ "step": 560
207
+ },
208
+ {
209
+ "epoch": 0.03909673070441524,
210
+ "grad_norm": 0.61328125,
211
+ "learning_rate": 0.0007,
212
+ "loss": 3.2046344757080076,
213
+ "step": 580
214
+ },
215
+ {
216
+ "epoch": 0.04044489383215369,
217
+ "grad_norm": 0.5234375,
218
+ "learning_rate": 0.0007,
219
+ "loss": 3.17300968170166,
220
+ "step": 600
221
+ },
222
+ {
223
+ "epoch": 0.04179305695989215,
224
+ "grad_norm": 0.578125,
225
+ "learning_rate": 0.0007,
226
+ "loss": 3.1147342681884767,
227
+ "step": 620
228
+ },
229
+ {
230
+ "epoch": 0.043141220087630605,
231
+ "grad_norm": 0.4765625,
232
+ "learning_rate": 0.0007,
233
+ "loss": 3.0973543167114257,
234
+ "step": 640
235
+ },
236
+ {
237
+ "epoch": 0.044489383215369056,
238
+ "grad_norm": 0.453125,
239
+ "learning_rate": 0.0007,
240
+ "loss": 3.043070602416992,
241
+ "step": 660
242
+ },
243
+ {
244
+ "epoch": 0.045837546343107514,
245
+ "grad_norm": 0.455078125,
246
+ "learning_rate": 0.0007,
247
+ "loss": 3.0213077545166014,
248
+ "step": 680
249
+ },
250
+ {
251
+ "epoch": 0.04718570947084597,
252
+ "grad_norm": 0.419921875,
253
+ "learning_rate": 0.0007,
254
+ "loss": 2.977329063415527,
255
+ "step": 700
256
+ },
257
+ {
258
+ "epoch": 0.04853387259858443,
259
+ "grad_norm": 0.490234375,
260
+ "learning_rate": 0.0007,
261
+ "loss": 2.942974090576172,
262
+ "step": 720
263
+ },
264
+ {
265
+ "epoch": 0.04988203572632288,
266
+ "grad_norm": 0.48046875,
267
+ "learning_rate": 0.0007,
268
+ "loss": 2.9064186096191404,
269
+ "step": 740
270
+ },
271
+ {
272
+ "epoch": 0.05123019885406134,
273
+ "grad_norm": 0.498046875,
274
+ "learning_rate": 0.0007,
275
+ "loss": 2.890872001647949,
276
+ "step": 760
277
+ },
278
+ {
279
+ "epoch": 0.0525783619817998,
280
+ "grad_norm": 0.5,
281
+ "learning_rate": 0.0007,
282
+ "loss": 2.8978702545166017,
283
+ "step": 780
284
+ },
285
+ {
286
+ "epoch": 0.053926525109538256,
287
+ "grad_norm": 0.466796875,
288
+ "learning_rate": 0.0007,
289
+ "loss": 2.8463741302490235,
290
+ "step": 800
291
+ },
292
+ {
293
+ "epoch": 0.05527468823727671,
294
+ "grad_norm": 0.47265625,
295
+ "learning_rate": 0.0007,
296
+ "loss": 2.8323719024658205,
297
+ "step": 820
298
+ },
299
+ {
300
+ "epoch": 0.056622851365015166,
301
+ "grad_norm": 0.5,
302
+ "learning_rate": 0.0007,
303
+ "loss": 2.8007469177246094,
304
+ "step": 840
305
+ },
306
+ {
307
+ "epoch": 0.057971014492753624,
308
+ "grad_norm": 0.43359375,
309
+ "learning_rate": 0.0007,
310
+ "loss": 2.7573240280151365,
311
+ "step": 860
312
+ },
313
+ {
314
+ "epoch": 0.05931917762049208,
315
+ "grad_norm": 0.470703125,
316
+ "learning_rate": 0.0007,
317
+ "loss": 2.7341793060302733,
318
+ "step": 880
319
+ },
320
+ {
321
+ "epoch": 0.06066734074823053,
322
+ "grad_norm": 0.546875,
323
+ "learning_rate": 0.0007,
324
+ "loss": 2.7180759429931642,
325
+ "step": 900
326
+ },
327
+ {
328
+ "epoch": 0.06201550387596899,
329
+ "grad_norm": 0.5625,
330
+ "learning_rate": 0.0007,
331
+ "loss": 2.692394256591797,
332
+ "step": 920
333
+ },
334
+ {
335
+ "epoch": 0.06336366700370745,
336
+ "grad_norm": 0.443359375,
337
+ "learning_rate": 0.0007,
338
+ "loss": 2.678949546813965,
339
+ "step": 940
340
+ },
341
+ {
342
+ "epoch": 0.06471183013144591,
343
+ "grad_norm": 0.466796875,
344
+ "learning_rate": 0.0007,
345
+ "loss": 2.651218223571777,
346
+ "step": 960
347
+ },
348
+ {
349
+ "epoch": 0.06605999325918437,
350
+ "grad_norm": 0.43359375,
351
+ "learning_rate": 0.0007,
352
+ "loss": 2.632099723815918,
353
+ "step": 980
354
+ },
355
+ {
356
+ "epoch": 0.06740815638692282,
357
+ "grad_norm": 0.470703125,
358
+ "learning_rate": 0.0007,
359
+ "loss": 2.6305675506591797,
360
+ "step": 1000
361
+ },
362
+ {
363
+ "epoch": 0.06875631951466127,
364
+ "grad_norm": 0.462890625,
365
+ "learning_rate": 0.0007,
366
+ "loss": 2.5994905471801757,
367
+ "step": 1020
368
+ },
369
+ {
370
+ "epoch": 0.07010448264239973,
371
+ "grad_norm": 0.474609375,
372
+ "learning_rate": 0.0007,
373
+ "loss": 2.5834341049194336,
374
+ "step": 1040
375
+ },
376
+ {
377
+ "epoch": 0.07145264577013818,
378
+ "grad_norm": 0.498046875,
379
+ "learning_rate": 0.0007,
380
+ "loss": 2.567459487915039,
381
+ "step": 1060
382
+ },
383
+ {
384
+ "epoch": 0.07280080889787664,
385
+ "grad_norm": 0.462890625,
386
+ "learning_rate": 0.0007,
387
+ "loss": 2.54827823638916,
388
+ "step": 1080
389
+ },
390
+ {
391
+ "epoch": 0.0741489720256151,
392
+ "grad_norm": 0.478515625,
393
+ "learning_rate": 0.0007,
394
+ "loss": 2.5431196212768556,
395
+ "step": 1100
396
+ },
397
+ {
398
+ "epoch": 0.07549713515335356,
399
+ "grad_norm": 0.4453125,
400
+ "learning_rate": 0.0007,
401
+ "loss": 2.5049072265625,
402
+ "step": 1120
403
+ },
404
+ {
405
+ "epoch": 0.07684529828109202,
406
+ "grad_norm": 0.4765625,
407
+ "learning_rate": 0.0007,
408
+ "loss": 2.4941583633422852,
409
+ "step": 1140
410
+ },
411
+ {
412
+ "epoch": 0.07819346140883048,
413
+ "grad_norm": 0.46875,
414
+ "learning_rate": 0.0007,
415
+ "loss": 2.4895328521728515,
416
+ "step": 1160
417
+ },
418
+ {
419
+ "epoch": 0.07954162453656892,
420
+ "grad_norm": 0.48828125,
421
+ "learning_rate": 0.0007,
422
+ "loss": 2.4841281890869142,
423
+ "step": 1180
424
+ },
425
+ {
426
+ "epoch": 0.08088978766430738,
427
+ "grad_norm": 0.46484375,
428
+ "learning_rate": 0.0007,
429
+ "loss": 2.475014495849609,
430
+ "step": 1200
431
+ }
432
+ ],
433
+ "logging_steps": 20,
434
+ "max_steps": 2500,
435
+ "num_input_tokens_seen": 0,
436
+ "num_train_epochs": 1,
437
+ "save_steps": 100,
438
+ "stateful_callbacks": {
439
+ "TrainerControl": {
440
+ "args": {
441
+ "should_epoch_stop": false,
442
+ "should_evaluate": false,
443
+ "should_log": false,
444
+ "should_save": true,
445
+ "should_training_stop": false
446
+ },
447
+ "attributes": {}
448
+ }
449
+ },
450
+ "total_flos": 3871540563148800.0,
451
+ "train_batch_size": 64,
452
+ "trial_name": null,
453
+ "trial_params": null
454
+ }
zain/Activation/out/glu-relu-100L_run/checkpoint-1200/training_args.bin ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:13465899109c0aad1422356a6b7684d9c92b0cc10d091b5c3e65ca518136a315
3
+ size 4920
zain/Activation/out/glu-relu-100L_run/checkpoint-1300/config.json ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "activation": "relu",
3
+ "architectures": [
4
+ "TinyLlamaForCausalLM"
5
+ ],
6
+ "attention_bias": false,
7
+ "attention_dropout": 0.0,
8
+ "bos_token_id": 1,
9
+ "dtype": "bfloat16",
10
+ "eos_token_id": 2,
11
+ "head_dim": 32,
12
+ "hidden_act": "silu",
13
+ "hidden_size": 128,
14
+ "initializer_range": 0.02,
15
+ "intermediate_size": 256,
16
+ "max_position_embeddings": 512,
17
+ "mlp_bias": false,
18
+ "mlp_type": "glu",
19
+ "model_type": "tiny_llama",
20
+ "num_attention_heads": 4,
21
+ "num_hidden_layers": 100,
22
+ "num_key_value_heads": 4,
23
+ "pad_token_id": 0,
24
+ "pretraining_tp": 1,
25
+ "rms_norm_eps": 1e-06,
26
+ "rope_parameters": {
27
+ "rope_theta": 10000.0,
28
+ "rope_type": "default"
29
+ },
30
+ "tie_word_embeddings": true,
31
+ "tokenizer_name": "w-ahmad/tiny-stories-tokenizer",
32
+ "transformers_version": "5.16.0.dev0",
33
+ "use_cache": false,
34
+ "vocab_size": 4096,
35
+ "waleed_beta": 4
36
+ }
zain/Activation/out/glu-relu-100L_run/checkpoint-1300/model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:f2617c871f04951e4385de2d29cf2d1a8aed437c0daceb4f0b01bc50c90afc1c
3
+ size 33967272
zain/Activation/out/glu-relu-100L_run/checkpoint-1300/optimizer.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d581932904814148864aba52f1d72fce2fc3cf6701de8432ccb9712d63f21002
3
+ size 68504996
zain/Activation/out/glu-relu-100L_run/checkpoint-1300/rng_state.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:9d9cd6a0487226e5bd30d1846894c82af483733ab4381b75bae9c0745e05d405
3
+ size 14244
zain/Activation/out/glu-relu-100L_run/checkpoint-1300/scheduler.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:6166e4e25b157aeaf4765c80f04dab16e63f8ec3cd6ab3637d77b940e7dfcfe1
3
+ size 1064
zain/Activation/out/glu-relu-100L_run/checkpoint-1300/tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
zain/Activation/out/glu-relu-100L_run/checkpoint-1300/tokenizer_config.json ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_prefix_space": false,
3
+ "backend": "tokenizers",
4
+ "bos_token": "<|endoftext|>",
5
+ "eos_token": "<|endoftext|>",
6
+ "errors": "replace",
7
+ "is_local": false,
8
+ "local_files_only": false,
9
+ "model_max_length": 1024,
10
+ "pad_token": "<|endoftext|>",
11
+ "tokenizer_class": "GPT2Tokenizer",
12
+ "unk_token": "<|endoftext|>"
13
+ }
zain/Activation/out/glu-relu-100L_run/checkpoint-1300/trainer_state.json ADDED
@@ -0,0 +1,489 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "best_global_step": null,
3
+ "best_metric": null,
4
+ "best_model_checkpoint": null,
5
+ "epoch": 0.08763060330299967,
6
+ "eval_steps": 2498,
7
+ "global_step": 1300,
8
+ "is_hyper_param_search": false,
9
+ "is_local_process_zero": true,
10
+ "is_world_process_zero": true,
11
+ "log_history": [
12
+ {
13
+ "epoch": 0.0013481631277384564,
14
+ "grad_norm": 2.09375,
15
+ "learning_rate": 2.66e-05,
16
+ "loss": 8.298905181884766,
17
+ "step": 20
18
+ },
19
+ {
20
+ "epoch": 0.002696326255476913,
21
+ "grad_norm": 1.3359375,
22
+ "learning_rate": 5.46e-05,
23
+ "loss": 8.045733642578124,
24
+ "step": 40
25
+ },
26
+ {
27
+ "epoch": 0.004044489383215369,
28
+ "grad_norm": 1.2734375,
29
+ "learning_rate": 8.259999999999999e-05,
30
+ "loss": 7.839435577392578,
31
+ "step": 60
32
+ },
33
+ {
34
+ "epoch": 0.005392652510953826,
35
+ "grad_norm": 1.265625,
36
+ "learning_rate": 0.0001106,
37
+ "loss": 7.574758148193359,
38
+ "step": 80
39
+ },
40
+ {
41
+ "epoch": 0.006740815638692282,
42
+ "grad_norm": 1.1953125,
43
+ "learning_rate": 0.0001386,
44
+ "loss": 7.235897064208984,
45
+ "step": 100
46
+ },
47
+ {
48
+ "epoch": 0.008088978766430738,
49
+ "grad_norm": 1.140625,
50
+ "learning_rate": 0.00016659999999999998,
51
+ "loss": 6.862369537353516,
52
+ "step": 120
53
+ },
54
+ {
55
+ "epoch": 0.009437141894169195,
56
+ "grad_norm": 1.0234375,
57
+ "learning_rate": 0.00019460000000000001,
58
+ "loss": 6.508391571044922,
59
+ "step": 140
60
+ },
61
+ {
62
+ "epoch": 0.010785305021907651,
63
+ "grad_norm": 1.078125,
64
+ "learning_rate": 0.0002226,
65
+ "loss": 6.215537261962891,
66
+ "step": 160
67
+ },
68
+ {
69
+ "epoch": 0.012133468149646108,
70
+ "grad_norm": 0.76171875,
71
+ "learning_rate": 0.00025059999999999997,
72
+ "loss": 5.984259414672851,
73
+ "step": 180
74
+ },
75
+ {
76
+ "epoch": 0.013481631277384564,
77
+ "grad_norm": 0.82421875,
78
+ "learning_rate": 0.0002786,
79
+ "loss": 5.756989669799805,
80
+ "step": 200
81
+ },
82
+ {
83
+ "epoch": 0.01482979440512302,
84
+ "grad_norm": 1.171875,
85
+ "learning_rate": 0.00030659999999999997,
86
+ "loss": 5.531437683105469,
87
+ "step": 220
88
+ },
89
+ {
90
+ "epoch": 0.016177957532861477,
91
+ "grad_norm": 1.1171875,
92
+ "learning_rate": 0.0003346,
93
+ "loss": 5.294247817993164,
94
+ "step": 240
95
+ },
96
+ {
97
+ "epoch": 0.01752612066059993,
98
+ "grad_norm": 0.9453125,
99
+ "learning_rate": 0.00036260000000000003,
100
+ "loss": 5.055661392211914,
101
+ "step": 260
102
+ },
103
+ {
104
+ "epoch": 0.01887428378833839,
105
+ "grad_norm": 0.94921875,
106
+ "learning_rate": 0.0003906,
107
+ "loss": 4.797793197631836,
108
+ "step": 280
109
+ },
110
+ {
111
+ "epoch": 0.020222446916076844,
112
+ "grad_norm": 1.421875,
113
+ "learning_rate": 0.0004186,
114
+ "loss": 4.594622802734375,
115
+ "step": 300
116
+ },
117
+ {
118
+ "epoch": 0.021570610043815303,
119
+ "grad_norm": 1.2265625,
120
+ "learning_rate": 0.0004466,
121
+ "loss": 4.39941635131836,
122
+ "step": 320
123
+ },
124
+ {
125
+ "epoch": 0.022918773171553757,
126
+ "grad_norm": 0.5546875,
127
+ "learning_rate": 0.00047460000000000004,
128
+ "loss": 4.24058723449707,
129
+ "step": 340
130
+ },
131
+ {
132
+ "epoch": 0.024266936299292215,
133
+ "grad_norm": 0.83984375,
134
+ "learning_rate": 0.0005026,
135
+ "loss": 4.102418899536133,
136
+ "step": 360
137
+ },
138
+ {
139
+ "epoch": 0.02561509942703067,
140
+ "grad_norm": 0.70703125,
141
+ "learning_rate": 0.0005306,
142
+ "loss": 3.975328063964844,
143
+ "step": 380
144
+ },
145
+ {
146
+ "epoch": 0.026963262554769128,
147
+ "grad_norm": 0.6171875,
148
+ "learning_rate": 0.0005586,
149
+ "loss": 3.839755630493164,
150
+ "step": 400
151
+ },
152
+ {
153
+ "epoch": 0.028311425682507583,
154
+ "grad_norm": 0.91796875,
155
+ "learning_rate": 0.0005866,
156
+ "loss": 3.7486804962158202,
157
+ "step": 420
158
+ },
159
+ {
160
+ "epoch": 0.02965958881024604,
161
+ "grad_norm": 0.61328125,
162
+ "learning_rate": 0.0006146,
163
+ "loss": 3.6831081390380858,
164
+ "step": 440
165
+ },
166
+ {
167
+ "epoch": 0.031007751937984496,
168
+ "grad_norm": 0.875,
169
+ "learning_rate": 0.0006426,
170
+ "loss": 3.5874557495117188,
171
+ "step": 460
172
+ },
173
+ {
174
+ "epoch": 0.032355915065722954,
175
+ "grad_norm": 0.546875,
176
+ "learning_rate": 0.0006705999999999999,
177
+ "loss": 3.5174354553222655,
178
+ "step": 480
179
+ },
180
+ {
181
+ "epoch": 0.03370407819346141,
182
+ "grad_norm": 0.66796875,
183
+ "learning_rate": 0.0006986,
184
+ "loss": 3.439300537109375,
185
+ "step": 500
186
+ },
187
+ {
188
+ "epoch": 0.03505224132119986,
189
+ "grad_norm": 0.54296875,
190
+ "learning_rate": 0.0007,
191
+ "loss": 3.386067581176758,
192
+ "step": 520
193
+ },
194
+ {
195
+ "epoch": 0.03640040444893832,
196
+ "grad_norm": 0.6015625,
197
+ "learning_rate": 0.0007,
198
+ "loss": 3.3094718933105467,
199
+ "step": 540
200
+ },
201
+ {
202
+ "epoch": 0.03774856757667678,
203
+ "grad_norm": 0.55859375,
204
+ "learning_rate": 0.0007,
205
+ "loss": 3.2609573364257813,
206
+ "step": 560
207
+ },
208
+ {
209
+ "epoch": 0.03909673070441524,
210
+ "grad_norm": 0.61328125,
211
+ "learning_rate": 0.0007,
212
+ "loss": 3.2046344757080076,
213
+ "step": 580
214
+ },
215
+ {
216
+ "epoch": 0.04044489383215369,
217
+ "grad_norm": 0.5234375,
218
+ "learning_rate": 0.0007,
219
+ "loss": 3.17300968170166,
220
+ "step": 600
221
+ },
222
+ {
223
+ "epoch": 0.04179305695989215,
224
+ "grad_norm": 0.578125,
225
+ "learning_rate": 0.0007,
226
+ "loss": 3.1147342681884767,
227
+ "step": 620
228
+ },
229
+ {
230
+ "epoch": 0.043141220087630605,
231
+ "grad_norm": 0.4765625,
232
+ "learning_rate": 0.0007,
233
+ "loss": 3.0973543167114257,
234
+ "step": 640
235
+ },
236
+ {
237
+ "epoch": 0.044489383215369056,
238
+ "grad_norm": 0.453125,
239
+ "learning_rate": 0.0007,
240
+ "loss": 3.043070602416992,
241
+ "step": 660
242
+ },
243
+ {
244
+ "epoch": 0.045837546343107514,
245
+ "grad_norm": 0.455078125,
246
+ "learning_rate": 0.0007,
247
+ "loss": 3.0213077545166014,
248
+ "step": 680
249
+ },
250
+ {
251
+ "epoch": 0.04718570947084597,
252
+ "grad_norm": 0.419921875,
253
+ "learning_rate": 0.0007,
254
+ "loss": 2.977329063415527,
255
+ "step": 700
256
+ },
257
+ {
258
+ "epoch": 0.04853387259858443,
259
+ "grad_norm": 0.490234375,
260
+ "learning_rate": 0.0007,
261
+ "loss": 2.942974090576172,
262
+ "step": 720
263
+ },
264
+ {
265
+ "epoch": 0.04988203572632288,
266
+ "grad_norm": 0.48046875,
267
+ "learning_rate": 0.0007,
268
+ "loss": 2.9064186096191404,
269
+ "step": 740
270
+ },
271
+ {
272
+ "epoch": 0.05123019885406134,
273
+ "grad_norm": 0.498046875,
274
+ "learning_rate": 0.0007,
275
+ "loss": 2.890872001647949,
276
+ "step": 760
277
+ },
278
+ {
279
+ "epoch": 0.0525783619817998,
280
+ "grad_norm": 0.5,
281
+ "learning_rate": 0.0007,
282
+ "loss": 2.8978702545166017,
283
+ "step": 780
284
+ },
285
+ {
286
+ "epoch": 0.053926525109538256,
287
+ "grad_norm": 0.466796875,
288
+ "learning_rate": 0.0007,
289
+ "loss": 2.8463741302490235,
290
+ "step": 800
291
+ },
292
+ {
293
+ "epoch": 0.05527468823727671,
294
+ "grad_norm": 0.47265625,
295
+ "learning_rate": 0.0007,
296
+ "loss": 2.8323719024658205,
297
+ "step": 820
298
+ },
299
+ {
300
+ "epoch": 0.056622851365015166,
301
+ "grad_norm": 0.5,
302
+ "learning_rate": 0.0007,
303
+ "loss": 2.8007469177246094,
304
+ "step": 840
305
+ },
306
+ {
307
+ "epoch": 0.057971014492753624,
308
+ "grad_norm": 0.43359375,
309
+ "learning_rate": 0.0007,
310
+ "loss": 2.7573240280151365,
311
+ "step": 860
312
+ },
313
+ {
314
+ "epoch": 0.05931917762049208,
315
+ "grad_norm": 0.470703125,
316
+ "learning_rate": 0.0007,
317
+ "loss": 2.7341793060302733,
318
+ "step": 880
319
+ },
320
+ {
321
+ "epoch": 0.06066734074823053,
322
+ "grad_norm": 0.546875,
323
+ "learning_rate": 0.0007,
324
+ "loss": 2.7180759429931642,
325
+ "step": 900
326
+ },
327
+ {
328
+ "epoch": 0.06201550387596899,
329
+ "grad_norm": 0.5625,
330
+ "learning_rate": 0.0007,
331
+ "loss": 2.692394256591797,
332
+ "step": 920
333
+ },
334
+ {
335
+ "epoch": 0.06336366700370745,
336
+ "grad_norm": 0.443359375,
337
+ "learning_rate": 0.0007,
338
+ "loss": 2.678949546813965,
339
+ "step": 940
340
+ },
341
+ {
342
+ "epoch": 0.06471183013144591,
343
+ "grad_norm": 0.466796875,
344
+ "learning_rate": 0.0007,
345
+ "loss": 2.651218223571777,
346
+ "step": 960
347
+ },
348
+ {
349
+ "epoch": 0.06605999325918437,
350
+ "grad_norm": 0.43359375,
351
+ "learning_rate": 0.0007,
352
+ "loss": 2.632099723815918,
353
+ "step": 980
354
+ },
355
+ {
356
+ "epoch": 0.06740815638692282,
357
+ "grad_norm": 0.470703125,
358
+ "learning_rate": 0.0007,
359
+ "loss": 2.6305675506591797,
360
+ "step": 1000
361
+ },
362
+ {
363
+ "epoch": 0.06875631951466127,
364
+ "grad_norm": 0.462890625,
365
+ "learning_rate": 0.0007,
366
+ "loss": 2.5994905471801757,
367
+ "step": 1020
368
+ },
369
+ {
370
+ "epoch": 0.07010448264239973,
371
+ "grad_norm": 0.474609375,
372
+ "learning_rate": 0.0007,
373
+ "loss": 2.5834341049194336,
374
+ "step": 1040
375
+ },
376
+ {
377
+ "epoch": 0.07145264577013818,
378
+ "grad_norm": 0.498046875,
379
+ "learning_rate": 0.0007,
380
+ "loss": 2.567459487915039,
381
+ "step": 1060
382
+ },
383
+ {
384
+ "epoch": 0.07280080889787664,
385
+ "grad_norm": 0.462890625,
386
+ "learning_rate": 0.0007,
387
+ "loss": 2.54827823638916,
388
+ "step": 1080
389
+ },
390
+ {
391
+ "epoch": 0.0741489720256151,
392
+ "grad_norm": 0.478515625,
393
+ "learning_rate": 0.0007,
394
+ "loss": 2.5431196212768556,
395
+ "step": 1100
396
+ },
397
+ {
398
+ "epoch": 0.07549713515335356,
399
+ "grad_norm": 0.4453125,
400
+ "learning_rate": 0.0007,
401
+ "loss": 2.5049072265625,
402
+ "step": 1120
403
+ },
404
+ {
405
+ "epoch": 0.07684529828109202,
406
+ "grad_norm": 0.4765625,
407
+ "learning_rate": 0.0007,
408
+ "loss": 2.4941583633422852,
409
+ "step": 1140
410
+ },
411
+ {
412
+ "epoch": 0.07819346140883048,
413
+ "grad_norm": 0.46875,
414
+ "learning_rate": 0.0007,
415
+ "loss": 2.4895328521728515,
416
+ "step": 1160
417
+ },
418
+ {
419
+ "epoch": 0.07954162453656892,
420
+ "grad_norm": 0.48828125,
421
+ "learning_rate": 0.0007,
422
+ "loss": 2.4841281890869142,
423
+ "step": 1180
424
+ },
425
+ {
426
+ "epoch": 0.08088978766430738,
427
+ "grad_norm": 0.46484375,
428
+ "learning_rate": 0.0007,
429
+ "loss": 2.475014495849609,
430
+ "step": 1200
431
+ },
432
+ {
433
+ "epoch": 0.08223795079204584,
434
+ "grad_norm": 0.43359375,
435
+ "learning_rate": 0.0007,
436
+ "loss": 2.459096908569336,
437
+ "step": 1220
438
+ },
439
+ {
440
+ "epoch": 0.0835861139197843,
441
+ "grad_norm": 0.443359375,
442
+ "learning_rate": 0.0007,
443
+ "loss": 2.4455087661743162,
444
+ "step": 1240
445
+ },
446
+ {
447
+ "epoch": 0.08493427704752275,
448
+ "grad_norm": 0.48828125,
449
+ "learning_rate": 0.0007,
450
+ "loss": 2.432474899291992,
451
+ "step": 1260
452
+ },
453
+ {
454
+ "epoch": 0.08628244017526121,
455
+ "grad_norm": 0.427734375,
456
+ "learning_rate": 0.0007,
457
+ "loss": 2.4141357421875,
458
+ "step": 1280
459
+ },
460
+ {
461
+ "epoch": 0.08763060330299967,
462
+ "grad_norm": 0.46875,
463
+ "learning_rate": 0.0007,
464
+ "loss": 2.4043970108032227,
465
+ "step": 1300
466
+ }
467
+ ],
468
+ "logging_steps": 20,
469
+ "max_steps": 2500,
470
+ "num_input_tokens_seen": 0,
471
+ "num_train_epochs": 1,
472
+ "save_steps": 100,
473
+ "stateful_callbacks": {
474
+ "TrainerControl": {
475
+ "args": {
476
+ "should_epoch_stop": false,
477
+ "should_evaluate": false,
478
+ "should_log": false,
479
+ "should_save": true,
480
+ "should_training_stop": false
481
+ },
482
+ "attributes": {}
483
+ }
484
+ },
485
+ "total_flos": 4194168943411200.0,
486
+ "train_batch_size": 64,
487
+ "trial_name": null,
488
+ "trial_params": null
489
+ }
zain/Activation/out/glu-relu-100L_run/checkpoint-1300/training_args.bin ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:13465899109c0aad1422356a6b7684d9c92b0cc10d091b5c3e65ca518136a315
3
+ size 4920
zain/Activation/out/glu-relu-100L_run/checkpoint-1400/config.json ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "activation": "relu",
3
+ "architectures": [
4
+ "TinyLlamaForCausalLM"
5
+ ],
6
+ "attention_bias": false,
7
+ "attention_dropout": 0.0,
8
+ "bos_token_id": 1,
9
+ "dtype": "bfloat16",
10
+ "eos_token_id": 2,
11
+ "head_dim": 32,
12
+ "hidden_act": "silu",
13
+ "hidden_size": 128,
14
+ "initializer_range": 0.02,
15
+ "intermediate_size": 256,
16
+ "max_position_embeddings": 512,
17
+ "mlp_bias": false,
18
+ "mlp_type": "glu",
19
+ "model_type": "tiny_llama",
20
+ "num_attention_heads": 4,
21
+ "num_hidden_layers": 100,
22
+ "num_key_value_heads": 4,
23
+ "pad_token_id": 0,
24
+ "pretraining_tp": 1,
25
+ "rms_norm_eps": 1e-06,
26
+ "rope_parameters": {
27
+ "rope_theta": 10000.0,
28
+ "rope_type": "default"
29
+ },
30
+ "tie_word_embeddings": true,
31
+ "tokenizer_name": "w-ahmad/tiny-stories-tokenizer",
32
+ "transformers_version": "5.16.0.dev0",
33
+ "use_cache": false,
34
+ "vocab_size": 4096,
35
+ "waleed_beta": 4
36
+ }
zain/Activation/out/glu-relu-100L_run/checkpoint-1400/model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:5e6a2130ba498fa0759d572b63b4f7071347adf63a7fe93f18b2c5c8e3e8e18b
3
+ size 33967272