path to 🎤MW/, fix bug

This commit is contained in:
billwuhao
2025-03-25 18:48:41 +08:00
parent e6a3707c1f
commit 9ca7a7cbca
19 changed files with 10 additions and 56 deletions
+2 -48
View File
@@ -122,7 +122,7 @@ class MultiLinePrompt:
},
}
CATEGORY = "MW/MW-DiffRhythm"
CATEGORY = "🎤MW/MW-DiffRhythm"
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("prompt",)
FUNCTION = "promptgen"
@@ -166,7 +166,7 @@ class DiffRhythmRun:
},
}
CATEGORY = "MW/MW-DiffRhythm"
CATEGORY = "🎤MW/MW-DiffRhythm"
RETURN_TYPES = ("AUDIO",)
RETURN_NAMES = ("audio",)
FUNCTION = "diffrhythmgen"
@@ -295,52 +295,6 @@ class DiffRhythmRun:
return text_emb
# @torch.no_grad()
# def get_style_prompt(self, model, audio=None, prompt=None):
# mulan = model
# if prompt is not None:
# return mulan(texts=prompt).half()
# if audio is None:
# raise ValueError("Audio data or style prompt must be provided")
# waveform = audio["waveform"]
# sample_rate = audio["sample_rate"]
# # Ensure waveform has correct shape
# if len(waveform.shape) == 3: # [1, channels, samples]
# waveform = waveform.squeeze(0)
# if waveform.shape[0] > 1: # If stereo, convert to mono
# waveform = waveform.mean(0, keepdim=True)
# # Calculate audio length (seconds)
# audio_len = waveform.shape[-1] / sample_rate
# if audio_len < 10:
# raise ValueError(f"Audio too short ({audio_len:.2f}s), minimum 10 seconds required.")
# # Extract middle 10-second segment
# mid_time = int((audio_len // 2) * sample_rate)
# start_sample = mid_time - int(5 * sample_rate)
# end_sample = start_sample + int(10 * sample_rate)
# wav_segment = waveform[..., start_sample:end_sample]
# # Resample to 24kHz
# if sample_rate != 24000:
# wav_segment = torchaudio.transforms.Resample(sample_rate, 24000)(wav_segment)
# # Ensure correct shape and device
# wav = wav_segment.to(model.device)
# if len(wav.shape) == 1:
# wav = wav.unsqueeze(0)
# with torch.no_grad():
# audio_emb = mulan(wavs=wav) # [1, 512]
# audio_emb = audio_emb.half()
# return audio_emb
def prepare_model(self, model, device, unload_model=False):
# prepare tokenizer
+5 -6
View File
@@ -17,7 +17,6 @@ class AudioRecorderDR:
"record_sec": ("INT", {
"default": 5,
"min": 1,
"max": 60,
"step": 1
}),
"sample_rate": (["16000", "44100", "48000"], {
@@ -31,14 +30,14 @@ class AudioRecorderDR:
}),
"sensitivity": ("FLOAT", {
"default": 1.2,
"min": 0.5,
"min": 0.1,
"max": 3.0,
"step": 0.1
}),
"smooth": ("INT", {
"default": 5,
"min": 1,
"max": 11,
"default": 1,
"min": 5,
"max": 7,
"step": 2
}),
"seed": ("INT", {
@@ -52,7 +51,7 @@ class AudioRecorderDR:
RETURN_TYPES = ("AUDIO",)
RETURN_NAMES = ("audio",)
FUNCTION = "record_and_clean"
CATEGORY = "MW/MW-DiffRhythm"
CATEGORY = "🎤MW/MW-DiffRhythm"
def _stft(self, y, n_fft):
hop = n_fft // 4
+2 -1
View File
@@ -31,7 +31,8 @@ from torch.optim.lr_scheduler import LinearLR, SequentialLR, ConstantLR
from accelerate import Accelerator
from accelerate.utils import DistributedDataParallelKwargs
from dataset.dataset import DiffusionDataset
from dr_dataset.dataset import DiffusionDataset
from torch.utils.data import DataLoader
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "diffrhythm_mw"
description = "Blazingly Fast and Embarrassingly Simple End-to-End Full-Length Song Generation. A node for ComfyUI."
version = "2.1.0"
version = "2.1.1"
license = {file = "LICENSE"}
dependencies = ["# accelerate==1.4.0", "# torchdiffeq==0.2.5", "# torchaudio==2.6.0", "# transformers==4.49.0", "# librosa==0.10.2.post1", "# pyarrow==19.0.1", "# pandas==2.2.3", "# bitsandbytes", "# jieba==0.42.1", "# cn2an==0.5.23", "# pypinyin==0.53.0", "# onnxruntime", "LangSegment", "x-transformers", "pylance", "ema-pytorch", "prefigure", "muq", "mutagen", "pyopenjtalk", "pykakasi", "Unidecode", "phonemizer"]