MCPcopy Create free account
hub / github.com/ace-step/ACE-Step / main

Function main

trainer.py:820–863  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

818
819
820def main(args):
821 model = Pipeline(
822 learning_rate=args.learning_rate,
823 num_workers=args.num_workers,
824 shift=args.shift,
825 max_steps=args.max_steps,
826 every_plot_step=args.every_plot_step,
827 dataset_path=args.dataset_path,
828 checkpoint_dir=args.checkpoint_dir,
829 adapter_name=args.exp_name,
830 lora_config_path=args.lora_config_path
831 )
832 checkpoint_callback = ModelCheckpoint(
833 monitor=None,
834 every_n_train_steps=args.every_n_train_steps,
835 save_top_k=-1,
836 )
837 # add datetime str to version
838 logger_callback = TensorBoardLogger(
839 version=datetime.now().strftime("%Y-%m-%d_%H-%M-%S") + args.exp_name,
840 save_dir=args.logger_dir,
841 )
842 trainer = Trainer(
843 accelerator="gpu",
844 devices=args.devices,
845 num_nodes=args.num_nodes,
846 precision=args.precision,
847 accumulate_grad_batches=args.accumulate_grad_batches,
848 strategy="ddp_find_unused_parameters_true",
849 max_epochs=args.epochs,
850 max_steps=args.max_steps,
851 log_every_n_steps=1,
852 logger=logger_callback,
853 callbacks=[checkpoint_callback],
854 gradient_clip_val=args.gradient_clip_val,
855 gradient_clip_algorithm=args.gradient_clip_algorithm,
856 reload_dataloaders_every_n_epochs=args.reload_dataloaders_every_n_epochs,
857 val_check_interval=args.val_check_interval,
858 )
859
860 trainer.fit(
861 model,
862 ckpt_path=args.ckpt_path,
863 )
864
865
866if __name__ == "__main__":

Callers 1

trainer.pyFile · 0.70

Calls 1

PipelineClass · 0.85

Tested by

no test coverage detected