From b1be0c7b082ddd98e47030ecb91b057022155fc6 Mon Sep 17 00:00:00 2001 From: holonic Date: Wed, 5 Jun 2024 22:30:43 +0100 Subject: [PATCH] removed __pycache__ --- .gitignore | 1 + __init__.py | 11 +++++ nodes.py | 104 ++++++++++++++++++++++++++++++++++++++ requirements.txt | 1 + util_config.py | 126 +++++++++++++++++++++++++++++++++++++++++++++++ 5 files changed, 243 insertions(+) create mode 100644 .gitignore create mode 100644 __init__.py create mode 100644 nodes.py create mode 100644 requirements.txt create mode 100644 util_config.py diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..bee8a64 --- /dev/null +++ b/.gitignore @@ -0,0 +1 @@ +__pycache__ diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..5f3c625 --- /dev/null +++ b/__init__.py @@ -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'] + diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..9361598 --- /dev/null +++ b/nodes.py @@ -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", +} diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..cb96583 --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +stable-audio-tools \ No newline at end of file diff --git a/util_config.py b/util_config.py new file mode 100644 index 0000000..13c5f77 --- /dev/null +++ b/util_config.py @@ -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] + } + } + }