Cion-lab commited on
Commit
87df298
·
verified ·
1 Parent(s): 7a667d0

preflight: stage pargs -- construct TrainingArguments/model/schedule on the Kaggle image, cross-check 22L=106,194,240 and 20L=99,114,048

Browse files
Files changed (1) hide show
  1. train/preflight.py +96 -3
train/preflight.py CHANGED
@@ -49,6 +49,97 @@ def last_json(text, begin, end):
49
  return None
50
 
51
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
52
  # ------------------------------------------------------------------------ P0 reader (CPU, free)
53
  def p0_reader(args, R):
54
  """The reader is where an unrecoverable main run would hide: if two runs at the same cursor read
@@ -219,7 +310,7 @@ def p3_cold_resume(args, R):
219
  def main():
220
  ap = argparse.ArgumentParser()
221
  ap.add_argument("--stage", required=True,
222
- choices=["p0", "p1", "p2", "p3", "summary", "all"])
223
  ap.add_argument("--seq-len", type=int, default=2048)
224
  ap.add_argument("--accum", type=int, default=32)
225
  ap.add_argument("--tokens", type=int, default=1_000_000_000)
@@ -235,8 +326,10 @@ def main():
235
  R = json.load(open(args.out_json))
236
  except Exception:
237
  print("existing preflight.json unreadable; starting fresh", flush=True)
238
- todo = ["p0", "p1", "p2", "p3"] if args.stage == "all" else [args.stage]
239
  t0 = time.time()
 
 
240
  if "p0" in todo:
241
  p0_reader(args, R)
242
  if "p1" in todo:
@@ -247,7 +340,7 @@ def main():
247
  p3_cold_resume(args, R)
248
  R["seconds_this_invocation"] = round(time.time() - t0, 1)
249
  R["gpu_hours_this_invocation"] = round(R["seconds_this_invocation"] / 3600.0, 3)
250
- R["PASSES"] = {k: R.get(k + "_pass") for k in ("P0", "P1", "P2", "P3")}
251
  R["GATE_3_READY"] = all(v is True for v in R["PASSES"].values())
252
  with open(args.out_json, "w") as f:
253
  json.dump(R, f, indent=1, default=str)
 
49
  return None
50
 
51
 
52
+ # ------------------------------------------------------------------------ Pargs (CPU, free)
53
+ def p_args(args, R):
54
+ """Every keyword in train_ounce100m.py's TrainingArguments is a promise about a library version, and
55
+ E-008/E-010 already taught us that transformers 5 drops things without warning. This constructs the
56
+ real objects on the Kaggle image -- TrainingArguments, the config, the model, the optimizer and the
57
+ trapezoid schedule -- with no data and no GPU, so a kwargs skew costs a free CPU minute instead of
58
+ killing a billed 2xT4 session at import time."""
59
+ src = r'''
60
+ import json, os, sys, math, time
61
+ sys.path.insert(0, "/kaggle/working")
62
+ os.environ.setdefault("CUDA_VISIBLE_DEVICES", "")
63
+ import torch
64
+ import transformers
65
+ from torch.optim import AdamW
66
+ 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)
76
+ 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,
83
+ adam_beta2=0.95, adam_epsilon=1e-8, max_grad_norm=1.0, lr_scheduler_type="constant",
84
+ warmup_ratio=0.0, num_train_epochs=1, max_steps=10, fp16=True, bf16=False,
85
+ gradient_checkpointing=False, ddp_find_unused_parameters=False, dataloader_num_workers=2,
86
+ dataloader_pin_memory=False, remove_unused_columns=False, ignore_data_skip=True,
87
+ save_strategy="steps", save_steps=5, save_total_limit=1, logging_steps=5, report_to=[],
88
+ seed=1, data_seed=1, accelerator_config={"dispatch_batches": False},
89
+ average_tokens_across_devices=False)
90
+ try:
91
+ ta = TrainingArguments(**kw); R["kwargs"] = "accepted"
92
+ except TypeError as e:
93
+ R["kwargs"] = f"REJECTED: {str(e)[:300]}"
94
+ bad = [k for k in kw if k in str(e)]
95
+ R["suspect_keys"] = bad
96
+ for k in bad:
97
+ kw.pop(k, None)
98
+ ta = TrainingArguments(**kw)
99
+ R["fp16_bf16"] = [ta.fp16, ta.bf16]
100
+ # schedule shape, standalone: warm 2%, flat to 80%, linear to zero
101
+ steps = 3815
102
+ opt = AdamW(m.parameters(), lr=6e-4, betas=(0.9, 0.95), weight_decay=0.1)
103
+ tr = T.TrapezoidTrainer.__new__(T.TrapezoidTrainer) # no Trainer.__init__: we want the schedule only
104
+ tr.lr_shape = {"warmup_steps": max(100, int(steps * 0.02)), "decay_start": int(steps * 0.8),
105
+ "total_steps": steps}
106
+ tr.optimizer, tr.lr_scheduler = opt, None
107
+ sch = tr.create_scheduler(steps, opt)
108
+ mults = []
109
+ for i in range(steps):
110
+ mults.append(sch.lr_lambdas[0](i))
111
+ sch.optimizer.param_groups[0]["lr"] = 6e-4 * mults[-1]
112
+ R["lr_shape"] = {"peak": max(mults), "at_warm_end": mults[76], "at_80pct": mults[int(steps*0.8)],
113
+ "at_90pct": mults[int(steps*0.9)], "final": mults[-1],
114
+ "plateau_fraction_steps": round(sum(1 for x in mults if x > 0.999) / steps, 3)}
115
+ # one real forward/backward on CPU with a tiny model to prove the collator output feeds the model
116
+ cfg2 = T.build_config(64, layers=2, hidden=128)
117
+ m2 = T.LlamaForCausalLM(cfg2)
118
+ store_ids = torch.randint(0, 1000, (3, 65))
119
+ out = m2(input_ids=store_ids[:, :-1], labels=store_ids[:, 1:],
120
+ attention_mask=torch.ones_like(store_ids[:, :-1]))
121
+ out.loss.backward()
122
+ R["forward_backward"] = {"loss": round(float(out.loss), 4), "finite": math.isfinite(float(out.loss)),
123
+ "grads": sum(1 for p in m2.parameters() if p.grad is not None)}
124
+ print("ARGS_JSON_BEGIN")
125
+ print(json.dumps(R, default=str))
126
+ print("ARGS_JSON_END")
127
+ '''
128
+ r = sh([sys.executable, "-c", src], timeout=3600, label="Pargs construct on the Kaggle image")
129
+ R["PARGS"] = last_json(r["out"], "ARGS_JSON_BEGIN", "ARGS_JSON_END")
130
+ R["PARGS_rc"] = r["rc"]
131
+ j = R["PARGS"] or {}
132
+ R["PARGS_pass"] = bool(r["rc"] == 0 and j.get("kwargs") == "accepted"
133
+ and j.get("params", {}).get("sum_numel") == 106194240
134
+ # 20L must be 99,114,048 per docs/01-plan.md §2.3's closed-form table, and its
135
+ # FFN must stay 1536; a doubled FFN would make test T1 meaningless.
136
+ and j.get("shapes_20L", {}).get("sum_numel") == 99114048
137
+ and j.get("shape20_ffn") == 1536
138
+ and (j.get("lr_shape") or {}).get("final") == 0.0
139
+ and (j.get("forward_backward") or {}).get("finite"))
140
+ print("VERDICT PARGS_pass=", R["PARGS_pass"], json.dumps(j.get("kwargs"))[:200], flush=True)
141
+
142
+
143
  # ------------------------------------------------------------------------ P0 reader (CPU, free)
144
  def p0_reader(args, R):
145
  """The reader is where an unrecoverable main run would hide: if two runs at the same cursor read
 
310
  def main():
311
  ap = argparse.ArgumentParser()
312
  ap.add_argument("--stage", required=True,
313
+ choices=["pargs", "p0", "p1", "p2", "p3", "summary", "all"])
314
  ap.add_argument("--seq-len", type=int, default=2048)
315
  ap.add_argument("--accum", type=int, default=32)
316
  ap.add_argument("--tokens", type=int, default=1_000_000_000)
 
326
  R = json.load(open(args.out_json))
327
  except Exception:
328
  print("existing preflight.json unreadable; starting fresh", flush=True)
329
+ todo = ["pargs", "p0", "p1", "p2", "p3"] if args.stage == "all" else [args.stage]
330
  t0 = time.time()
331
+ if "pargs" in todo:
332
+ p_args(args, R)
333
  if "p0" in todo:
334
  p0_reader(args, R)
335
  if "p1" in todo:
 
340
  p3_cold_resume(args, R)
341
  R["seconds_this_invocation"] = round(time.time() - t0, 1)
342
  R["gpu_hours_this_invocation"] = round(R["seconds_this_invocation"] / 3600.0, 3)
343
+ R["PASSES"] = {k: R.get(k + "_pass") for k in ("PARGS", "P0", "P1", "P2", "P3")}
344
  R["GATE_3_READY"] = all(v is True for v in R["PASSES"].values())
345
  with open(args.out_json, "w") as f:
346
  json.dump(R, f, indent=1, default=str)