Taimwe commited on
Commit
e9659f3
·
verified ·
1 Parent(s): 7c6176b

Option B: --mix-scale, trackio-before-config fix

Browse files
Files changed (1) hide show
  1. train_securecoder.py +10 -1
train_securecoder.py CHANGED
@@ -687,6 +687,7 @@ def init_trackio(args):
687
  log.info("trackio initialised (project=%s space=%s)", args.trackio_project, args.trackio_space)
688
  except Exception as exc: # noqa: BLE001 - monitoring must never kill training
689
  log.warning("trackio init failed (%s); continuing without it", exc)
 
690
 
691
 
692
  def train(args, model, tokenizer, records: list[dict]):
@@ -700,8 +701,8 @@ def train(args, model, tokenizer, records: list[dict]):
700
  log.info("train rows=%d eval rows=%d", len(train_ds), len(eval_ds) if eval_ds else 0)
701
 
702
  steps_per_epoch = len(train_ds) // max(args.batch_size * args.grad_accum, 1)
703
- cfg = build_sft_config(args, eval_ds is not None, steps_per_epoch)
704
  init_trackio(args)
 
705
 
706
  trainer = SFTTrainer(
707
  model=model,
@@ -787,6 +788,8 @@ def parse_args(argv=None):
787
  p.add_argument("--learning-rate", type=float, default=2e-4)
788
  p.add_argument("--lr-scheduler", default="cosine")
789
  p.add_argument("--num-epochs", type=float, default=1.0)
 
 
790
  p.add_argument("--max-steps", type=int, default=0, help="overrides --num-epochs when > 0")
791
  p.add_argument("--eval-samples", type=int, default=200, help="0 disables evaluation")
792
  p.add_argument("--logging-steps", type=int, default=10)
@@ -825,9 +828,15 @@ def apply_smoke(args) -> None:
825
 
826
 
827
  def main(argv=None) -> int:
 
828
  args = parse_args(argv)
829
  if args.smoke:
830
  apply_smoke(args)
 
 
 
 
 
831
  token = os.environ.get("HF_TOKEN")
832
 
833
  if args.validate_only:
 
687
  log.info("trackio initialised (project=%s space=%s)", args.trackio_project, args.trackio_space)
688
  except Exception as exc: # noqa: BLE001 - monitoring must never kill training
689
  log.warning("trackio init failed (%s); continuing without it", exc)
690
+ args.report_to = "none"
691
 
692
 
693
  def train(args, model, tokenizer, records: list[dict]):
 
701
  log.info("train rows=%d eval rows=%d", len(train_ds), len(eval_ds) if eval_ds else 0)
702
 
703
  steps_per_epoch = len(train_ds) // max(args.batch_size * args.grad_accum, 1)
 
704
  init_trackio(args)
705
+ cfg = build_sft_config(args, eval_ds is not None, steps_per_epoch)
706
 
707
  trainer = SFTTrainer(
708
  model=model,
 
788
  p.add_argument("--learning-rate", type=float, default=2e-4)
789
  p.add_argument("--lr-scheduler", default="cosine")
790
  p.add_argument("--num-epochs", type=float, default=1.0)
791
+ p.add_argument("--mix-scale", type=float, default=1.0,
792
+ help="multiply every source limit by this (e.g. 0.5 for a half mix)")
793
  p.add_argument("--max-steps", type=int, default=0, help="overrides --num-epochs when > 0")
794
  p.add_argument("--eval-samples", type=int, default=200, help="0 disables evaluation")
795
  p.add_argument("--logging-steps", type=int, default=10)
 
828
 
829
 
830
  def main(argv=None) -> int:
831
+ global MIX
832
  args = parse_args(argv)
833
  if args.smoke:
834
  apply_smoke(args)
835
+ elif args.mix_scale != 1.0:
836
+ MIX = [Source(s.repo, max(50, int(s.limit * args.mix_scale)),
837
+ s.kind, s.config, s.split, s.note) for s in MIX]
838
+ log.info("mix scaled by %.2f -> %d rows planned", args.mix_scale,
839
+ sum(s.limit for s in MIX))
840
  token = os.environ.get("HF_TOKEN")
841
 
842
  if args.validate_only: