fix train bugs and add train details
This commit is contained in:
+44
-15
@@ -64,7 +64,7 @@ class Pipeline(LightningModule):
|
||||
|
||||
# step 1: load model
|
||||
acestep_pipeline = ACEStepPipeline(checkpoint_dir)
|
||||
acestep_pipeline.load_checkpoint(checkpoint_dir)
|
||||
acestep_pipeline.load_checkpoint(acestep_pipeline.checkpoint_dir)
|
||||
|
||||
transformers = acestep_pipeline.ace_step_transformer.float().cpu()
|
||||
|
||||
@@ -90,9 +90,33 @@ class Pipeline(LightningModule):
|
||||
if self.is_train:
|
||||
self.transformers.train()
|
||||
|
||||
self.mert_model = AutoModel.from_pretrained(
|
||||
"m-a-p/MERT-v1-330M", trust_remote_code=True, cache_dir=checkpoint_dir
|
||||
).eval()
|
||||
# download first
|
||||
try:
|
||||
self.mert_model = AutoModel.from_pretrained(
|
||||
"m-a-p/MERT-v1-330M", trust_remote_code=True, cache_dir=checkpoint_dir
|
||||
).eval()
|
||||
except:
|
||||
import json
|
||||
import os
|
||||
|
||||
mert_config_path = os.path.join(
|
||||
os.path.expanduser("~"),
|
||||
".cache",
|
||||
"huggingface",
|
||||
"hub",
|
||||
"models--m-a-p--MERT-v1-330M",
|
||||
"blobs",
|
||||
"14f770758c7fe5c5e8ead4fe0f8e5fa727eb6942"
|
||||
)
|
||||
|
||||
with open(mert_config_path) as f:
|
||||
mert_config = json.load(f)
|
||||
mert_config["conv_pos_batch_norm"] = False
|
||||
with open(mert_config_path, mode="w") as f:
|
||||
json.dump(mert_config, f)
|
||||
self.mert_model = AutoModel.from_pretrained(
|
||||
"m-a-p/MERT-v1-330M", trust_remote_code=True, cache_dir=checkpoint_dir
|
||||
).eval()
|
||||
self.mert_model.requires_grad_(False)
|
||||
self.resampler_mert = torchaudio.transforms.Resample(
|
||||
orig_freq=48000, new_freq=24000
|
||||
@@ -101,18 +125,13 @@ class Pipeline(LightningModule):
|
||||
"m-a-p/MERT-v1-330M", trust_remote_code=True
|
||||
)
|
||||
|
||||
self.hubert_model = AutoModel.from_pretrained(
|
||||
"utter-project/mHuBERT-147",
|
||||
local_files_only=True,
|
||||
cache_dir=checkpoint_dir,
|
||||
).eval()
|
||||
self.hubert_model = AutoModel.from_pretrained("utter-project/mHuBERT-147").eval()
|
||||
self.hubert_model.requires_grad_(False)
|
||||
self.resampler_mhubert = torchaudio.transforms.Resample(
|
||||
orig_freq=48000, new_freq=16000
|
||||
)
|
||||
self.processor_mhubert = Wav2Vec2FeatureExtractor.from_pretrained(
|
||||
"utter-project/mHuBERT-147",
|
||||
local_files_only=True,
|
||||
cache_dir=checkpoint_dir,
|
||||
)
|
||||
|
||||
@@ -578,6 +597,16 @@ class Pipeline(LightningModule):
|
||||
def training_step(self, batch, batch_idx):
|
||||
return self.run_step(batch, batch_idx)
|
||||
|
||||
def on_save_checkpoint(self, checkpoint):
|
||||
state = {}
|
||||
log_dir = self.logger.log_dir
|
||||
epoch = self.current_epoch
|
||||
step = self.global_step
|
||||
checkpoint_name = f"epoch={epoch}-step={step}_lora"
|
||||
checkpoint_dir = os.path.join(log_dir, "checkpoints", checkpoint_name)
|
||||
self.transformers.save_lora_adapter(checkpoint_dir, adapter_name=self.adapter_name)
|
||||
return state
|
||||
|
||||
@torch.no_grad()
|
||||
def diffusion_process(
|
||||
self,
|
||||
@@ -808,7 +837,7 @@ def main(args):
|
||||
num_nodes=args.num_nodes,
|
||||
precision=args.precision,
|
||||
accumulate_grad_batches=args.accumulate_grad_batches,
|
||||
strategy="deepspeed_stage_2",
|
||||
strategy="ddp_find_unused_parameters_true",
|
||||
max_epochs=args.epochs,
|
||||
max_steps=args.max_steps,
|
||||
log_every_n_steps=1,
|
||||
@@ -835,9 +864,9 @@ if __name__ == "__main__":
|
||||
args.add_argument("--epochs", type=int, default=-1)
|
||||
args.add_argument("--max_steps", type=int, default=2000000)
|
||||
args.add_argument("--every_n_train_steps", type=int, default=2000)
|
||||
args.add_argument("--dataset_path", type=str, default="./data/your_dataset_path")
|
||||
args.add_argument("--exp_name", type=str, default="text2music_train_test")
|
||||
args.add_argument("--precision", type=str, default="bf16-mixed")
|
||||
args.add_argument("--dataset_path", type=str, default="./zh_lora_dataset")
|
||||
args.add_argument("--exp_name", type=str, default="chinese_rap_lora")
|
||||
args.add_argument("--precision", type=str, default="32")
|
||||
args.add_argument("--accumulate_grad_batches", type=int, default=1)
|
||||
args.add_argument("--devices", type=int, default=1)
|
||||
args.add_argument("--logger_dir", type=str, default="./exps/logs/")
|
||||
@@ -848,6 +877,6 @@ if __name__ == "__main__":
|
||||
args.add_argument("--reload_dataloaders_every_n_epochs", type=int, default=1)
|
||||
args.add_argument("--every_plot_step", type=int, default=2000)
|
||||
args.add_argument("--val_check_interval", type=int, default=None)
|
||||
args.add_argument("--lora_config_path", type=str, default=None)
|
||||
args.add_argument("--lora_config_path", type=str, default="config/zh_rap_lora_config.json")
|
||||
args = args.parse_args()
|
||||
main(args)
|
||||
|
||||
Reference in New Issue
Block a user