From 6e6dc1730f8f6d4267c774b6b45e6573d9662162 Mon Sep 17 00:00:00 2001 From: Gong Junmin <1836678486@qq.com> Date: Thu, 15 May 2025 22:39:41 +0800 Subject: [PATCH] Revert "fix_train_bug:adapter_name" --- trainer.py | 4 ---- 1 file changed, 4 deletions(-) diff --git a/trainer.py b/trainer.py index fce7cd7..fdc2efd 100644 --- a/trainer.py +++ b/trainer.py @@ -52,7 +52,6 @@ class Pipeline(LightningModule): warmup_steps: int = 4000, dataset_path: str = "./data/your_dataset_path", lora_config_path: str = None, - adapter_name: str = "lora_adapter", ): super().__init__() @@ -69,7 +68,6 @@ class Pipeline(LightningModule): transformers = acestep_pipeline.ace_step_transformer.float().cpu() - assert lora_config_path is not None, "Please provide a LoRA config path" if lora_config_path is not None: try: from peft import LoraConfig @@ -79,7 +77,6 @@ class Pipeline(LightningModule): lora_config = json.load(f) lora_config = LoraConfig(**lora_config) transformers.add_adapter(adapter_config=lora_config) - self.adapter_name = adapter_name self.transformers = transformers @@ -824,7 +821,6 @@ def main(args): every_plot_step=args.every_plot_step, dataset_path=args.dataset_path, checkpoint_dir=args.checkpoint_dir, - adapter_name=args.exp_name, ) checkpoint_callback = ModelCheckpoint( monitor=None,