Robust Hub push + throughput reporting
Browse files- 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 |
-
|
| 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 |
-
|
| 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}")
|