Push adapter at every checkpoint (--push-checkpoints) so cancel/timeout cannot lose a run
Browse files- train_securecoder.py +51 -0
train_securecoder.py
CHANGED
|
@@ -802,6 +802,32 @@ def train(args, model, tokenizer, records: list[dict]):
|
|
| 802 |
)
|
| 803 |
|
| 804 |
started = time.time()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 805 |
stats = trainer.train()
|
| 806 |
elapsed = time.time() - started
|
| 807 |
log.info("training finished in %.1f min (final loss %.4f)",
|
|
@@ -819,6 +845,28 @@ def train(args, model, tokenizer, records: list[dict]):
|
|
| 819 |
return trainer, stats, elapsed
|
| 820 |
|
| 821 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 822 |
def _push_adapter(api, model, args):
|
| 823 |
"""Unsloth's push_to_hub signature varies between releases (one version
|
| 824 |
rejects ``tokenizer=``), so fall back to uploading the saved folder - the
|
|
@@ -867,6 +915,9 @@ def parse_args(argv=None):
|
|
| 867 |
|
| 868 |
p.add_argument("--base-model", default="Qwen/Qwen3-Coder-30B-A3B-Instruct")
|
| 869 |
p.add_argument("--output-repo", default=None, help="Hub repo for the LoRA adapter")
|
|
|
|
|
|
|
|
|
|
| 870 |
p.add_argument("--resume-adapter", default=None,
|
| 871 |
help="existing LoRA repo to keep training (does not init a fresh adapter)")
|
| 872 |
p.add_argument("--merge-repo", default=None, help="optional Hub repo for a 16-bit merge")
|
|
|
|
| 802 |
)
|
| 803 |
|
| 804 |
started = time.time()
|
| 805 |
+
|
| 806 |
+
if args.push_checkpoints and args.output_repo:
|
| 807 |
+
# Push the *adapter* at every checkpoint save, so a cancel or a timeout
|
| 808 |
+
# still leaves a resumable adapter on the hub. Repository creation races
|
| 809 |
+
# with `save_and_push` below, hence exist_ok=True.
|
| 810 |
+
from huggingface_hub import HfApi
|
| 811 |
+
from transformers import TrainerCallback
|
| 812 |
+
|
| 813 |
+
repo_id, api = args.output_repo, HfApi()
|
| 814 |
+
if os.environ.get("HF_TOKEN"):
|
| 815 |
+
api.create_repo(repo_id, repo_type="model", private=False, exist_ok=True)
|
| 816 |
+
|
| 817 |
+
class _PushAdapterOnSave(TrainerCallback):
|
| 818 |
+
def on_save(self, _args, state, control, **kwargs):
|
| 819 |
+
if state.global_step <= 0:
|
| 820 |
+
return
|
| 821 |
+
try:
|
| 822 |
+
_upload_adapter_files(api, f"checkpoints/checkpoint-{state.global_step}",
|
| 823 |
+
repo_id)
|
| 824 |
+
log.info("checkpoint adapter @step %d pushed to %s (%.1f min in)",
|
| 825 |
+
state.global_step, repo_id, (time.time() - started) / 60)
|
| 826 |
+
except Exception as exc: # noqa: BLE001 - never break training
|
| 827 |
+
log.warning("checkpoint push failed (continuing): %s", exc)
|
| 828 |
+
|
| 829 |
+
trainer.add_callback(_PushAdapterOnSave())
|
| 830 |
+
|
| 831 |
stats = trainer.train()
|
| 832 |
elapsed = time.time() - started
|
| 833 |
log.info("training finished in %.1f min (final loss %.4f)",
|
|
|
|
| 845 |
return trainer, stats, elapsed
|
| 846 |
|
| 847 |
|
| 848 |
+
def _upload_adapter_files(api, folder: str, repo_id: str) -> None:
|
| 849 |
+
"""Upload just the LoRA adapter (and any tokenizer files) from ``folder``
|
| 850 |
+
to the repo root, so the repo always holds the latest *resumable* adapter.
|
| 851 |
+
|
| 852 |
+
A canceled or timed-out job then leaves behind something ``--resume-adapter``
|
| 853 |
+
can continue from, instead of losing the whole run.
|
| 854 |
+
"""
|
| 855 |
+
try:
|
| 856 |
+
names = os.listdir(folder)
|
| 857 |
+
except OSError:
|
| 858 |
+
return
|
| 859 |
+
for name in sorted(names):
|
| 860 |
+
wanted = (name.startswith(("adapter_", "tokenizer", "chat_template"))
|
| 861 |
+
or name in ("special_tokens_map.json", "added_tokens.json"))
|
| 862 |
+
if not wanted:
|
| 863 |
+
continue
|
| 864 |
+
path = os.path.join(folder, name)
|
| 865 |
+
if os.path.isfile(path):
|
| 866 |
+
api.upload_file(path_or_fileobj=path, path_in_repo=name,
|
| 867 |
+
repo_id=repo_id, repo_type="model")
|
| 868 |
+
|
| 869 |
+
|
| 870 |
def _push_adapter(api, model, args):
|
| 871 |
"""Unsloth's push_to_hub signature varies between releases (one version
|
| 872 |
rejects ``tokenizer=``), so fall back to uploading the saved folder - the
|
|
|
|
| 915 |
|
| 916 |
p.add_argument("--base-model", default="Qwen/Qwen3-Coder-30B-A3B-Instruct")
|
| 917 |
p.add_argument("--output-repo", default=None, help="Hub repo for the LoRA adapter")
|
| 918 |
+
p.add_argument("--push-checkpoints", action="store_true",
|
| 919 |
+
help="push the adapter to --output-repo at every save step, so a "
|
| 920 |
+
"canceled/timed-out job still leaves a resumable adapter")
|
| 921 |
p.add_argument("--resume-adapter", default=None,
|
| 922 |
help="existing LoRA repo to keep training (does not init a fresh adapter)")
|
| 923 |
p.add_argument("--merge-repo", default=None, help="optional Hub repo for a 16-bit merge")
|