Taimwe commited on
Commit
5bf8c15
·
verified ·
1 Parent(s): 884b7d7

Push adapter at every checkpoint (--push-checkpoints) so cancel/timeout cannot lose a run

Browse files
Files changed (1) hide show
  1. 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")