Cion-lab commited on
Commit
658b0a8
·
verified ·
1 Parent(s): 11456a3

train: preflight -- static-review fixes: absolute-step cursor, hub_listing pagination, cursor required on resume, repo_type everywhere, rank-0 repo create

Browse files
Files changed (1) hide show
  1. train/preflight.py +20 -7
train/preflight.py CHANGED
@@ -67,9 +67,13 @@ from transformers import TrainingArguments
67
  import train_ounce100m as T
68
  R = {"transformers": transformers.__version__, "torch": torch.__version__}
69
  cfg = T.build_config(2048)
70
- R["config"] = {k: getattr(cfg, k) for k in ("hidden_size", "num_hidden_layers", "num_attention_heads",
71
- "num_key_value_heads", "intermediate_size", "vocab_size",
72
- "tie_word_embeddings", "max_position_embeddings")}
 
 
 
 
73
  t0 = time.time()
74
  m = T.LlamaForCausalLM(cfg)
75
  R["params"] = T.count_params(m)
@@ -77,6 +81,12 @@ R["init_seconds"] = round(time.time() - t0, 1)
77
  R["shapes_20L"] = T.count_params(T.LlamaForCausalLM(T.build_config(1024, layers=20)))
78
  c20 = T.build_config(1024, layers=20)
79
  R["shape20_ffn"] = c20.intermediate_size
 
 
 
 
 
 
80
  # The exact kwargs the trainer passes. If the image rejects one, this is where we learn it.
81
  kw = dict(output_dir="/kaggle/working/argcheck", per_device_train_batch_size=2,
82
  gradient_accumulation_steps=32, learning_rate=6e-4, weight_decay=0.1, adam_beta1=0.9,
@@ -161,6 +171,7 @@ _emit(R)
161
  # FFN must stay 1536; a doubled FFN would make test T1 meaningless.
162
  and j.get("shapes_20L", {}).get("sum_numel") == 99114048
163
  and j.get("shape20_ffn") == 1536 and shape_ok
 
164
  and (j.get("forward_backward") or {}).get("finite"))
165
  print("VERDICT PARGS_pass=", R["PARGS_pass"], "blocks=", R["PARGS_n_blocks"], "shape_ok=", shape_ok,
166
  json.dumps({k: j.get(k) for k in ("transformers", "torch", "kwargs", "params", "shapes_20L",
@@ -294,9 +305,10 @@ print("HUB_JSON_END")
294
 
295
  # ------------------------------------------------------------------------ P2..P4 (GPU)
296
  def torchrun(args, extra, timeout=7200):
297
- return sh(["torchrun", "--nproc_per_node=2", "train_ounce100m.py", "--root",
298
- WORK + "/mixroot", "--seq-len", str(args.seq_len)] + extra,
299
- timeout=timeout, label="torchrun " + " ".join(extra[:6]))
 
300
 
301
 
302
  def p2_throughput(args, R):
@@ -349,7 +361,8 @@ def p3_cold_resume(args, R):
349
  """Short run -> wipe every local trace -> resume with --resume auto, which must recover the exact
350
  sample position from the Hub. Loss continuity is the pass condition: a restart that silently
351
  re-initialised would jump the loss."""
352
- common = ["--tokens", str(args.tokens), "--accum", str(args.accum), "--micro-batch", "2",
 
353
  "--hub-repo", args.ckpt_repo, "--prune", "--log-every", "5"]
354
  a = torchrun(args, common + ["--max-steps", str(args.steps_a), "--out", WORK + "/p3"],
355
  timeout=9000)
 
67
  import train_ounce100m as T
68
  R = {"transformers": transformers.__version__, "torch": torch.__version__}
69
  cfg = T.build_config(2048)
70
+ R["config"] = {k: getattr(cfg, k, "<absent>") for k in
71
+ ("hidden_size", "num_hidden_layers", "num_attention_heads", "num_key_value_heads",
72
+ "intermediate_size", "vocab_size", "tie_word_embeddings", "max_position_embeddings",
73
+ # rope_theta moved into rope_parameters in newer transformers; if the flat attribute is
74
+ # gone the value we think we froze may not be the value the model got (verify, not assume)
75
+ "rope_theta", "rope_parameters", "rms_norm_eps", "attention_dropout", "mlp_bias",
76
+ "hidden_act")}
77
  t0 = time.time()
78
  m = T.LlamaForCausalLM(cfg)
79
  R["params"] = T.count_params(m)
 
81
  R["shapes_20L"] = T.count_params(T.LlamaForCausalLM(T.build_config(1024, layers=20)))
82
  c20 = T.build_config(1024, layers=20)
83
  R["shape20_ffn"] = c20.intermediate_size
84
+ R["shape20_kv"] = c20.num_key_value_heads
85
+ # Regression test for the defect that would have made T1 meaningless: passing --hidden 576 explicitly is
86
+ # the frozen width, so it must NOT fall through to the variant branch and become MHA-9 with a doubled FFN.
87
+ cv = T.build_config(2048, hidden=576)
88
+ R["hidden576_still_frozen"] = (cv.num_key_value_heads == 3 and cv.intermediate_size == 1536
89
+ and cv.num_hidden_layers == 22)
90
  # The exact kwargs the trainer passes. If the image rejects one, this is where we learn it.
91
  kw = dict(output_dir="/kaggle/working/argcheck", per_device_train_batch_size=2,
92
  gradient_accumulation_steps=32, learning_rate=6e-4, weight_decay=0.1, adam_beta1=0.9,
 
171
  # FFN must stay 1536; a doubled FFN would make test T1 meaningless.
172
  and j.get("shapes_20L", {}).get("sum_numel") == 99114048
173
  and j.get("shape20_ffn") == 1536 and shape_ok
174
+ and j.get("hidden576_still_frozen") is True
175
  and (j.get("forward_backward") or {}).get("finite"))
176
  print("VERDICT PARGS_pass=", R["PARGS_pass"], "blocks=", R["PARGS_n_blocks"], "shape_ok=", shape_ok,
177
  json.dumps({k: j.get(k) for k in ("transformers", "torch", "kwargs", "params", "shapes_20L",
 
305
 
306
  # ------------------------------------------------------------------------ P2..P4 (GPU)
307
  def torchrun(args, extra, timeout=7200):
308
+ # --seq-len is deliberately NOT set here: each caller passes its own, and relying on argparse's
309
+ # last-wins to undo a duplicate is a defect waiting for a reordering.
310
+ return sh(["torchrun", "--nproc_per_node=2", "train_ounce100m.py", "--root", WORK + "/mixroot"]
311
+ + extra, timeout=timeout, label="torchrun " + " ".join(extra[:6]))
312
 
313
 
314
  def p2_throughput(args, R):
 
361
  """Short run -> wipe every local trace -> resume with --resume auto, which must recover the exact
362
  sample position from the Hub. Loss continuity is the pass condition: a restart that silently
363
  re-initialised would jump the loss."""
364
+ common = ["--seq-len", str(args.seq_len), "--tokens", str(args.tokens),
365
+ "--accum", str(args.accum), "--micro-batch", "2",
366
  "--hub-repo", args.ckpt_repo, "--prune", "--log-every", "5"]
367
  a = torchrun(args, common + ["--max-steps", str(args.steps_a), "--out", WORK + "/p3"],
368
  timeout=9000)