Keep attnres on resume, fix final chunk_index, default 0.5b to Yi-6B
CLI default attnres=off was overwriting block checkpoints on resume so later loads hit Unexpected key(s). Only apply flags the user passed. Track next_chunk so a budget-exit save does not skip the untrained yield. 0.5b now uses 01-ai/Yi-6B (64k); refuse resume when the ckpt tokenizer does not match.
This commit is contained in:
+39
-23
@@ -41,7 +41,7 @@ _TOY_TRAIN = {
|
||||
"gen_every": 200,
|
||||
}
|
||||
_B500M_TRAIN = {
|
||||
"tokenizer": "Qwen/Qwen3-8B",
|
||||
"tokenizer": "01-ai/Yi-6B",
|
||||
"out": "ckpts/k3_0.5b.pt",
|
||||
"limit": 20000,
|
||||
"batch": 2,
|
||||
@@ -184,6 +184,20 @@ def _apply_moe_coefs(model, cfg: K3Config) -> None:
|
||||
module.z_loss_coef = cfg.moe_z_loss_coef
|
||||
|
||||
|
||||
def _apply_cli_overrides(cfg: K3Config, args: argparse.Namespace) -> None:
|
||||
"""Copy only flags the user actually passed. CLI defaults must not clobber a resume."""
|
||||
if args.attnres is not None:
|
||||
cfg.attnres = args.attnres
|
||||
if args.attnres_block_size is not None:
|
||||
cfg.attnres_block_size = args.attnres_block_size
|
||||
if args.grad_checkpoint is not None:
|
||||
cfg.gradient_checkpointing = args.grad_checkpoint
|
||||
if args.moe_aux_coef is not None:
|
||||
cfg.moe_aux_loss_coef = args.moe_aux_coef
|
||||
if args.moe_z_coef is not None:
|
||||
cfg.moe_z_loss_coef = args.moe_z_coef
|
||||
|
||||
|
||||
def _moe_log(model) -> dict:
|
||||
frac = moe_route_frac(model)
|
||||
if frac is None:
|
||||
@@ -267,9 +281,10 @@ def main() -> None:
|
||||
p.add_argument("--device", default="auto")
|
||||
p.add_argument(
|
||||
"--attnres",
|
||||
default="off",
|
||||
default=None,
|
||||
choices=["off", "full", "block"],
|
||||
help="depth mixer: off=standard residual, block=K3 AttnRes, full=per-layer AttnRes",
|
||||
help="depth mixer: off=standard residual (preset default), block=K3 AttnRes, "
|
||||
"full=per-layer AttnRes. Omit on --resume to keep the checkpoint value",
|
||||
)
|
||||
p.add_argument(
|
||||
"--attnres-block-size",
|
||||
@@ -329,14 +344,7 @@ def main() -> None:
|
||||
tok = load_tokenizer(args.tokenizer)
|
||||
cfg = K3Config.preset(args.preset)
|
||||
cfg.vocab_size = tok.vocab_size
|
||||
cfg.attnres = args.attnres
|
||||
cfg.attnres_block_size = args.attnres_block_size
|
||||
if args.grad_checkpoint is not None:
|
||||
cfg.gradient_checkpointing = args.grad_checkpoint
|
||||
if args.moe_aux_coef is not None:
|
||||
cfg.moe_aux_loss_coef = args.moe_aux_coef
|
||||
if args.moe_z_coef is not None:
|
||||
cfg.moe_z_loss_coef = args.moe_z_coef
|
||||
_apply_cli_overrides(cfg, args)
|
||||
|
||||
langs = [part.strip() for part in args.langs.split(",") if part.strip()]
|
||||
tpm = tokens_per_micro(args.batch, args.seq_len)
|
||||
@@ -364,19 +372,16 @@ def main() -> None:
|
||||
f"{type(loaded_cfg).__name__}"
|
||||
)
|
||||
cfg = loaded_cfg
|
||||
cfg.attnres = args.attnres
|
||||
cfg.attnres_block_size = args.attnres_block_size
|
||||
if args.grad_checkpoint is not None:
|
||||
cfg.gradient_checkpointing = args.grad_checkpoint
|
||||
if args.moe_aux_coef is not None:
|
||||
cfg.moe_aux_loss_coef = args.moe_aux_coef
|
||||
if args.moe_z_coef is not None:
|
||||
cfg.moe_z_loss_coef = args.moe_z_coef
|
||||
_apply_cli_overrides(cfg, args)
|
||||
model.gradient_checkpointing = cfg.gradient_checkpointing
|
||||
model.to(device)
|
||||
payload = torch.load(args.resume, map_location="cpu", weights_only=False)
|
||||
if payload.get("tokenizer") and payload["tokenizer"] != args.tokenizer:
|
||||
print(f"warning: ckpt tokenizer {payload['tokenizer']} != {args.tokenizer}")
|
||||
raise SystemExit(
|
||||
f"tokenizer mismatch: ckpt {payload['tokenizer']!r} vs "
|
||||
f"CLI {args.tokenizer!r}; embeddings are not interchangeable "
|
||||
f"(do not resume a Qwen ckpt with Yi)"
|
||||
)
|
||||
micro_step = int(payload.get("micro_step", 0))
|
||||
opt_step = int(payload.get("opt_step", 0))
|
||||
tokens = int(payload.get("tokens", 0))
|
||||
@@ -469,6 +474,10 @@ def main() -> None:
|
||||
model.train()
|
||||
t0 = time.perf_counter()
|
||||
tokens_at_t0 = tokens
|
||||
# Index of the next untrained chunk. Mid-loop saves use last_trained+1.
|
||||
# The final save must NOT +1 again: the loop may break on a yielded chunk
|
||||
# that was never trained (budget check is at the top).
|
||||
next_chunk = chunk_index
|
||||
for chunk_index, x, y in iter_indexed(train_chunks, start=chunk_index):
|
||||
if args.max_tokens is not None:
|
||||
if tokens >= args.max_tokens:
|
||||
@@ -497,6 +506,7 @@ def main() -> None:
|
||||
lr_now = optim.param_groups[0]["lr"]
|
||||
if raw_loss < best_train:
|
||||
best_train = raw_loss
|
||||
next_chunk = chunk_index + 1
|
||||
|
||||
metrics = {
|
||||
"train/loss": raw_loss,
|
||||
@@ -542,7 +552,7 @@ def main() -> None:
|
||||
micro_step=micro_step,
|
||||
opt_step=opt_step,
|
||||
tokens=tokens,
|
||||
chunk_index=chunk_index + 1,
|
||||
chunk_index=next_chunk,
|
||||
best_heldout=best_heldout,
|
||||
)
|
||||
_save(_sibling(args.out, "_best"), payload)
|
||||
@@ -570,7 +580,7 @@ def main() -> None:
|
||||
micro_step=micro_step,
|
||||
opt_step=opt_step,
|
||||
tokens=tokens,
|
||||
chunk_index=chunk_index + 1,
|
||||
chunk_index=next_chunk,
|
||||
best_heldout=best_heldout,
|
||||
)
|
||||
_save(_sibling(args.out, "_last"), payload)
|
||||
@@ -586,7 +596,7 @@ def main() -> None:
|
||||
micro_step=micro_step,
|
||||
opt_step=opt_step,
|
||||
tokens=tokens,
|
||||
chunk_index=chunk_index + 1,
|
||||
chunk_index=next_chunk,
|
||||
best_heldout=best_heldout,
|
||||
)
|
||||
_save(args.out, payload)
|
||||
@@ -596,6 +606,12 @@ def main() -> None:
|
||||
)
|
||||
if tracker is not None:
|
||||
tracker.finish()
|
||||
if device == "cuda":
|
||||
try:
|
||||
torch.cuda.synchronize()
|
||||
torch.cuda.empty_cache()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
Reference in New Issue
Block a user