preflight: stage pargs -- construct TrainingArguments/model/schedule on the Kaggle image, cross-check 22L=106,194,240 and 20L=99,114,048
Browse files- 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)
|