path to 🎤MW/, fix bug
This commit is contained in:
+2
-48
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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"]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user