removed __pycache__
This commit is contained in:
@@ -0,0 +1 @@
|
||||
__pycache__
|
||||
+11
@@ -0,0 +1,11 @@
|
||||
"""
|
||||
@author: lks-ai
|
||||
@title: StableAudioSampler
|
||||
@nickname: stableaudio
|
||||
@description: A Simple integration of Stable Audio Diffusion with knobs and stuff!
|
||||
"""
|
||||
|
||||
from .nodes import StableAudioSampler, NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
import os
|
||||
import glob
|
||||
import torch
|
||||
import torchaudio
|
||||
from einops import rearrange
|
||||
from stable_audio_tools import get_pretrained_model
|
||||
from stable_audio_tools.inference.generation import generate_diffusion_cond
|
||||
|
||||
from safetensors.torch import load_file
|
||||
from .util_config import get_model_config
|
||||
from stable_audio_tools.models.factory import create_model_from_config
|
||||
from stable_audio_tools.models.utils import load_ckpt_state_dict
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
class AnyType(str):
|
||||
def __ne__(self, __value: object) -> bool:
|
||||
return False
|
||||
base_path = os.path.dirname(os.path.realpath(__file__))
|
||||
|
||||
# Our any instance wants to be a wildcard string
|
||||
any = AnyType("audio")
|
||||
model_files = [os.path.basename(file) for file in glob.glob("models/audio_checkpoints/*.safetensors")] + [os.path.basename(file) for file in glob.glob("models/audio_checkpoints/*.ckpt")]
|
||||
if len(model_files) == 0:
|
||||
model_files.append("Put models in models/audio_checkpoints")
|
||||
|
||||
|
||||
def generate_audio(prompt, steps, cfg_scale, sample_size, sigma_min, sigma_max, sampler_type, device, save, save_path, model_filename):
|
||||
model_path = f"models/audio_checkpoints/{model_filename}"
|
||||
if model_filename.endswith(".safetensors") or model_filename.endswith(".ckpt"):
|
||||
model = create_model_from_config(get_model_config())
|
||||
model.load_state_dict(load_ckpt_state_dict(model_path))
|
||||
else:
|
||||
model, model_config = get_pretrained_model("stabilityai/stable-audio-open-1.0")
|
||||
sample_rate = model_config["sample_rate"]
|
||||
sample_size = model_config["sample_size"]
|
||||
|
||||
model = model.to(device)
|
||||
|
||||
conditioning = [{
|
||||
"prompt": prompt,
|
||||
"seconds_start": 0,
|
||||
"seconds_total": 30
|
||||
}]
|
||||
|
||||
output = generate_diffusion_cond(
|
||||
model,
|
||||
steps=steps,
|
||||
cfg_scale=cfg_scale,
|
||||
conditioning=conditioning,
|
||||
sample_size=sample_size,
|
||||
sigma_min=sigma_min,
|
||||
sigma_max=sigma_max,
|
||||
sampler_type=sampler_type,
|
||||
device=device
|
||||
)
|
||||
|
||||
output = rearrange(output, "b d n -> d (b n)")
|
||||
|
||||
output = output.to(torch.float32).div(torch.max(torch.abs(output))).clamp(-1, 1).mul(32767).to(torch.int16).cpu()
|
||||
|
||||
if save:
|
||||
torchaudio.save("output/" + save_path, output, sample_rate)
|
||||
|
||||
# Convert to bytes
|
||||
audio_bytes = output.numpy().tobytes()
|
||||
|
||||
return audio_bytes, sample_rate
|
||||
|
||||
class StableAudioSampler:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING", {"default": "128 BPM tech house drum loop"}),
|
||||
"model_filename": (model_files, ),
|
||||
"steps": ("INT", {"default": 100, "min": 1, "max": 10000}),
|
||||
"cfg_scale": ("FLOAT", {"default": 7.0, "min": 0.0, "max": 100.0, "step": 0.1}),
|
||||
"sample_size": ("INT", {"default": 65536, "min": 1, "max": 1000000}),
|
||||
"sigma_min": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 1000.0, "step": 0.01}),
|
||||
"sigma_max": ("FLOAT", {"default": 500.0, "min": 0.0, "max": 1000.0, "step": 0.01}),
|
||||
"sampler_type": ("STRING", {"default": "dpmpp-3m-sde"}),
|
||||
"save": ("BOOLEAN", {"default": True}),
|
||||
"save_path": ("STRING", {"default": "output.wav"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = (any, "INT")
|
||||
FUNCTION = "sample"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
CATEGORY = "audio"
|
||||
|
||||
def sample(self, prompt, steps, cfg_scale, sample_size, sigma_min, sigma_max, sampler_type, save, save_path, model_filename):
|
||||
audio_bytes, sample_rate = generate_audio(prompt, steps, cfg_scale, sample_size, sigma_min, sigma_max, sampler_type, device, save, save_path, model_filename)
|
||||
return (audio_bytes, sample_rate)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"StableAudioSampler": StableAudioSampler,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"StableAudioSampler": "Stable Diffusion Audio Sampler",
|
||||
}
|
||||
@@ -0,0 +1 @@
|
||||
stable-audio-tools
|
||||
+126
@@ -0,0 +1,126 @@
|
||||
def get_model_config():
|
||||
return {
|
||||
"model_type": "diffusion_cond",
|
||||
"sample_size": 2097152,
|
||||
"sample_rate": 44100,
|
||||
"audio_channels": 2,
|
||||
"model": {
|
||||
"pretransform": {
|
||||
"type": "autoencoder",
|
||||
"iterate_batch": True,
|
||||
"config": {
|
||||
"encoder": {
|
||||
"type": "oobleck",
|
||||
"requires_grad": False,
|
||||
"config": {
|
||||
"in_channels": 2,
|
||||
"channels": 128,
|
||||
"c_mults": [1, 2, 4, 8, 16],
|
||||
"strides": [2, 4, 4, 8, 8],
|
||||
"latent_dim": 128,
|
||||
"use_snake": True
|
||||
}
|
||||
},
|
||||
"decoder": {
|
||||
"type": "oobleck",
|
||||
"config": {
|
||||
"out_channels": 2,
|
||||
"channels": 128,
|
||||
"c_mults": [1, 2, 4, 8, 16],
|
||||
"strides": [2, 4, 4, 8, 8],
|
||||
"latent_dim": 64,
|
||||
"use_snake": True,
|
||||
"final_tanh": False
|
||||
}
|
||||
},
|
||||
"bottleneck": {
|
||||
"type": "vae"
|
||||
},
|
||||
"latent_dim": 64,
|
||||
"downsampling_ratio": 2048,
|
||||
"io_channels": 2
|
||||
}
|
||||
},
|
||||
"conditioning": {
|
||||
"configs": [
|
||||
{
|
||||
"id": "prompt",
|
||||
"type": "t5",
|
||||
"config": {
|
||||
"t5_model_name": "t5-base",
|
||||
"max_length": 128
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": "seconds_start",
|
||||
"type": "number",
|
||||
"config": {
|
||||
"min_val": 0,
|
||||
"max_val": 512
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": "seconds_total",
|
||||
"type": "number",
|
||||
"config": {
|
||||
"min_val": 0,
|
||||
"max_val": 512
|
||||
}
|
||||
}
|
||||
],
|
||||
"cond_dim": 768
|
||||
},
|
||||
"diffusion": {
|
||||
"cross_attention_cond_ids": ["prompt", "seconds_start", "seconds_total"],
|
||||
"global_cond_ids": ["seconds_start", "seconds_total"],
|
||||
"type": "dit",
|
||||
"config": {
|
||||
"io_channels": 64,
|
||||
"embed_dim": 1536,
|
||||
"depth": 24,
|
||||
"num_heads": 24,
|
||||
"cond_token_dim": 768,
|
||||
"global_cond_dim": 1536,
|
||||
"project_cond_tokens": False,
|
||||
"transformer_type": "continuous_transformer"
|
||||
}
|
||||
},
|
||||
"io_channels": 64
|
||||
},
|
||||
"training": {
|
||||
"use_ema": True,
|
||||
"log_loss_info": False,
|
||||
"optimizer_configs": {
|
||||
"diffusion": {
|
||||
"optimizer": {
|
||||
"type": "AdamW",
|
||||
"config": {
|
||||
"lr": 5e-5,
|
||||
"betas": [0.9, 0.999],
|
||||
"weight_decay": 1e-3
|
||||
}
|
||||
},
|
||||
"scheduler": {
|
||||
"type": "InverseLR",
|
||||
"config": {
|
||||
"inv_gamma": 1000000,
|
||||
"power": 0.5,
|
||||
"warmup": 0.99
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"demo": {
|
||||
"demo_every": 2000,
|
||||
"demo_steps": 250,
|
||||
"num_demos": 4,
|
||||
"demo_cond": [
|
||||
{"prompt": "Amen break 174 BPM", "seconds_start": 0, "seconds_total": 12},
|
||||
{"prompt": "A beautiful orchestral symphony, classical music", "seconds_start": 0, "seconds_total": 160},
|
||||
{"prompt": "Chill hip-hop beat, chillhop", "seconds_start": 0, "seconds_total": 190},
|
||||
{"prompt": "A pop song about love and loss", "seconds_start": 0, "seconds_total": 180}
|
||||
],
|
||||
"demo_cfg_scales": [3, 6, 9]
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user