Taimwe commited on
Commit
4d01873
·
verified ·
1 Parent(s): 1e16000

Robust Hub push + throughput reporting

Browse files
Files changed (1) hide show
  1. train_securecoder.py +33 -2
train_securecoder.py CHANGED
@@ -729,6 +729,29 @@ def train(args, model, tokenizer, records: list[dict]):
729
  return trainer, stats, elapsed
730
 
731
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
732
  def save_and_push(args, model, tokenizer):
733
  from huggingface_hub import HfApi
734
 
@@ -737,13 +760,14 @@ def save_and_push(args, model, tokenizer):
737
 
738
  model.save_pretrained(args.output_dir)
739
  tokenizer.save_pretrained(args.output_dir)
 
740
  log.info("pushing LoRA adapter to %s", args.output_repo)
741
- model.push_to_hub(args.output_repo, tokenizer=tokenizer)
742
 
743
  if args.merge_repo:
744
  api.create_repo(args.merge_repo, repo_type="model", exist_ok=True, private=args.private)
745
  log.info("merging to 16-bit and pushing to %s (large upload)", args.merge_repo)
746
- model.push_to_hub_merged(args.merge_repo, tokenizer=tokenizer, save_method="merged_16bit")
747
 
748
  # --------------------------------------------------------------------------
749
  # CLI
@@ -845,6 +869,13 @@ def main(argv=None) -> int:
845
  print("\n" + "=" * 78)
846
  print(f"DONE rows={len(records):,} time={elapsed / 60:.1f} min "
847
  f"loss={stats_train.metrics.get('train_loss', float('nan')):.4f}")
 
 
 
 
 
 
 
848
  print(f"adapter: https://huggingface.co/{args.output_repo}")
849
  if args.merge_repo:
850
  print(f"merged : https://huggingface.co/{args.merge_repo}")
 
729
  return trainer, stats, elapsed
730
 
731
 
732
+ def _push_adapter(api, model, args):
733
+ """Unsloth's push_to_hub signature varies between releases (one version
734
+ rejects ``tokenizer=``), so fall back to uploading the saved folder - the
735
+ tokenizer files sit next to the adapter and get pushed either way."""
736
+ try:
737
+ model.push_to_hub(args.output_repo)
738
+ return
739
+ except TypeError as exc:
740
+ log.warning("push_to_hub rejected our arguments (%s); uploading the folder instead", exc)
741
+ except Exception as exc: # noqa: BLE001 - never lose a finished run to an upload quirk
742
+ log.warning("model.push_to_hub failed (%s); uploading the folder instead", exc)
743
+ api.upload_folder(folder_path=args.output_dir, repo_id=args.output_repo, repo_type="model")
744
+
745
+
746
+ def _push_merged(api, model, tokenizer, args):
747
+ try:
748
+ model.push_to_hub_merged(args.merge_repo, tokenizer=tokenizer, save_method="merged_16bit")
749
+ return
750
+ except TypeError as exc:
751
+ log.warning("push_to_hub_merged rejected tokenizer= (%s); retrying without it", exc)
752
+ model.push_to_hub_merged(args.merge_repo, save_method="merged_16bit")
753
+
754
+
755
  def save_and_push(args, model, tokenizer):
756
  from huggingface_hub import HfApi
757
 
 
760
 
761
  model.save_pretrained(args.output_dir)
762
  tokenizer.save_pretrained(args.output_dir)
763
+
764
  log.info("pushing LoRA adapter to %s", args.output_repo)
765
+ _push_adapter(api, model, args)
766
 
767
  if args.merge_repo:
768
  api.create_repo(args.merge_repo, repo_type="model", exist_ok=True, private=args.private)
769
  log.info("merging to 16-bit and pushing to %s (large upload)", args.merge_repo)
770
+ _push_merged(api, model, tokenizer, args)
771
 
772
  # --------------------------------------------------------------------------
773
  # CLI
 
869
  print("\n" + "=" * 78)
870
  print(f"DONE rows={len(records):,} time={elapsed / 60:.1f} min "
871
  f"loss={stats_train.metrics.get('train_loss', float('nan')):.4f}")
872
+ speed = stats_train.metrics.get("train_samples_per_second") or 0
873
+ if speed:
874
+ est_hours = len(records) / speed / 3600
875
+ print(f"throughput: {speed:.1f} rows/s -> a 1-epoch pass over {len(records):,} rows "
876
+ f"≈ {est_hours:.1f} h at this rate")
877
+ print(f"cost at $1.80/h (l40sx1): ≈ ${est_hours * 1.80:.0f} "
878
+ f"at $2.50/h (a100-large): ≈ ${est_hours * 2.50:.0f}")
879
  print(f"adapter: https://huggingface.co/{args.output_repo}")
880
  if args.merge_repo:
881
  print(f"merged : https://huggingface.co/{args.merge_repo}")