fix train bugs and add train details

This commit is contained in:
chuxij
2025-05-12 18:50:54 +00:00
parent 84ba6afea3
commit ce2caec957
9 changed files with 261 additions and 143 deletions
+44 -15
View File
@@ -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)