49 Commits
Author SHA1 Message Date
Salvador E. Tropea 621bd2784e Bumped version to 1.1.3 2026-02-11 11:03:35 -03:00
Salvador E. Tropea 4e56851264 [Fixed] Issues on Windows when auto-downloading
Fixes #4
2026-02-11 10:53:37 -03:00
Salvador E. Tropea 8a84f6b48e Bumped version to 1.1.2 2025-07-27 14:40:42 -03:00
Salvador E. Tropea 2737d7b190 [Added] Version information
- When registering the nodes
- To command line tools
2025-07-27 14:38:26 -03:00
Salvador E. Tropea 335af9fa4b [Added] Pre-commit script to check version consistency
ComfyUI registry fault, poor support of pyproject.toml options
2025-07-27 14:36:23 -03:00
Salvador E. Tropea accb0fd56c [DOCs][Added] README for the icons 2025-07-27 12:48:57 -03:00
Salvador E. Tropea 329a8826ad [CI/CD][Changed] To publish on tag (with semantic version) 2025-07-27 12:44:52 -03:00
Salvador E. Tropea d8f885e61a [DOCs] Updated installation and description
Also added Icon and Banner entries for ComfyUI manager
2025-07-27 12:40:20 -03:00
Salvador E. Tropea 82e0f78e21 Ignore sync script 2025-07-27 12:34:27 -03:00
Salvador E. Tropea 17de6ca853 [Torch] Moved to SeCoNoHe 2025-07-23 13:15:25 -03:00
Salvador E. Tropea 87aacaa903 [Tools] Adapted to SeCoNoHe 2025-07-23 10:44:11 -03:00
Salvador E. Tropea 4da4a7b498 [Logger] Now using get_debug_level and debugl from SeCoNoHe 2025-07-23 09:37:01 -03:00
Salvador E. Tropea f0cf78326c [Requirements][Added] SeCoNoHe 2025-07-23 09:31:14 -03:00
Salvador E. Tropea d387703101 Migrated to use SeCoNoHe 2025-07-23 09:20:16 -03:00
Salvador E. Tropea 1d4782516c Bumped version to 1.1.1 2025-07-21 09:39:57 -03:00
Salvador E. Tropea ad8ad49d9c [Examples][Added] Quick versions
They download the audio from internet
In most cases just uses 10 seconds to make it faster
2025-07-20 19:17:10 -03:00
Salvador E. Tropea 0b425c9bd1 [Demucs][Removed] Normalization
Not really needed
2025-07-18 11:03:54 -03:00
Salvador E. Tropea a377c72fac [Examples][Added] Demix and Remix example
Also added links to the examples from the README
2025-07-18 10:58:29 -03:00
Salvador E. Tropea ab3b9d1c3d [Demucs][Added] Torchaudio model support
The one you get using HDEMUCS_HIGH_MUSDB_PLUS
Is an HDemucs model, the code in TorchAudio looks old, perhaps
simplified with incompatible renames.
2025-07-12 13:49:32 -03:00
Salvador E. Tropea c87923faeb Bumped version to 1.1.0 2025-07-11 18:04:42 -03:00
Salvador E. Tropea 64070adf77 [Tool] Made all executable 2025-07-11 17:52:20 -03:00
Salvador E. Tropea b8ee1856c5 [Demix][Fixed] Handling of not generated stems 2025-07-11 17:51:45 -03:00
Salvador E. Tropea f973796f07 [Demucs][Replaced] Julius up/down sampler by torchaudio
Has less distortion.
2025-07-11 17:14:05 -03:00
Salvador E. Tropea 9a1297beaf [Demucs][Logger][Added] Signature
And made more robust the sync between sig and kwargs
2025-07-11 17:12:42 -03:00
Salvador E. Tropea 93856c88eb [Demucs][Logger][Fixed] Wiener iterations default 2025-07-11 17:12:02 -03:00
Salvador E. Tropea 75ed37ed3b [Fixed] Restored wiener
Used by the mdx.safetensors file, one HDemucs uses CaC and the
other Wiener
2025-07-11 17:09:35 -03:00
Salvador E. Tropea 08fb457a9d [Demucs][Removed] Weiner filtering support
Currently unused, will be restored only if we find a model using it
2025-07-11 10:01:33 -03:00
Salvador E. Tropea dfcce31318 [Demucs][Logger][Added] Better STFT framework log 2025-07-11 09:59:04 -03:00
Salvador E. Tropea 499aac5e80 [Demucs] Better logger
- Better code
- Better output
2025-07-11 09:24:53 -03:00
Salvador E. Tropea 6760256cd7 [Added] Debug information about Demucs
So we can know the architecture used
2025-07-10 13:08:18 -03:00
Salvador E. Tropea 19321451ce [Models DB][Fixed] Problems with already linked children 2025-07-08 11:07:45 -03:00
Salvador E. Tropea b44e229116 [Added] Support for partial models created from bigger ones
I guess they are to test individual portions
2025-07-08 10:51:16 -03:00
Salvador E. Tropea 7f51868c42 [DOCs][Added] How we load a model 2025-07-08 10:31:33 -03:00
Salvador E. Tropea 5bcecd9268 [Demix][Fix] Demucs segments
So now the default is valid for MDX and Demucs
2025-07-08 10:28:45 -03:00
Salvador E. Tropea b76546eb3e [Demix] Use canonical device names
To avoid the code confusing cuda vs cuda:0
2025-07-08 10:27:24 -03:00
Salvador E. Tropea 683525aa73 [Demucs to Safetensors] Run without demucs lib
Even when we didn't use the lib the torch.load needed it to map
the classes during the pickle load.
So now we remap the lib modules to our own copy of the classes.
2025-07-08 10:24:42 -03:00
Salvador E. Tropea 9fec04d0fb [Added] Ignore the folder where I store the Demucs parts 2025-07-06 15:11:14 -03:00
Salvador E. Tropea bd18e15833 [Added] Demucs workflow example 2025-07-06 15:10:12 -03:00
Salvador E. Tropea 2ca35cffd9 [Added] Demucs support and node 2025-07-06 15:09:36 -03:00
Salvador E. Tropea 0ae82e60b4 [Load Model][Added] Support for other models
Just avoid accessing data particular for a model
2025-07-06 15:07:40 -03:00
Salvador E. Tropea f1841269a0 [Load Safetensors][Added] Support for multiple models in one file
"Bag of Models" in Demucs tongue
2025-07-06 15:05:41 -03:00
Salvador E. Tropea e6d0a289cf [Logger][Removed] Unneeded line 2025-07-06 15:04:10 -03:00
Salvador E. Tropea 53dad57654 [Misc][Added] Helpers to serialize Fractions
Needed for Demucs
2025-07-06 15:02:58 -03:00
Salvador E. Tropea 1a01e4461c [Torch][Added] get_offload_device and get_canonical_device
The first to be in a neutral place and be optional.
The second allows better device name comparisson.
2025-07-06 15:01:30 -03:00
Salvador E. Tropea 8cbcb4518a [Requirements][Added] Missing safetensors 2025-07-06 15:00:25 -03:00
Salvador E. Tropea a953355e7d [Tools][Demix][Added] Support for optional stems 2025-07-06 14:54:34 -03:00
Salvador E. Tropea a34e05062e [Added] Demuc models we known to the database 2025-07-06 14:53:09 -03:00
Salvador E. Tropea 0dc28bcdd8 [Added] Tool to convert Demucs into safetensors and see its metadata 2025-07-06 14:52:04 -03:00
Salvador E. Tropea e7221bce0e [Added] Demucs code from Meta 2025-07-06 14:46:25 -03:00
87 changed files with 6115 additions and 1166 deletions
@@ -2,10 +2,8 @@ name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- main
paths:
- "pyproject.toml"
tags:
- '[0-9]+.[0-9]+.[0-9]+' # e.g. 1.2.3
jobs:
publish-node:
@@ -14,6 +12,7 @@ jobs:
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
with:
+3
View File
@@ -9,6 +9,9 @@ test/
.*~
*.mp3
models/all_public_uvr_models
models/demucs
__no__
models/.catalog.csv
models/*.yaml
0LEEME
sync.sh
+15
View File
@@ -41,3 +41,18 @@ repos:
# "--check-hidden"
]
# You can create a .codespellignore file with one word per line for words to ignore.
# --- Version checking ---
- repo: local
hooks:
- id: version-check
name: check for version consistency
# The command to execute. It's a Python script.
entry: python3 tool/check_versions.py
# Use 'system' to run it with the current environment's Python
language: system
# This hook doesn't need to run on every file.
# It should run if either of the version files change.
# This makes it very fast.
files: ^(pyproject\.toml|src/nodes/__init__\.py)$
# The regex `^...$` ensures it matches the full path from the repo root.
+97 -24
View File
@@ -14,12 +14,14 @@ audio demixing, also known as audio separation.
From an audio the objective is to separate the vocals, instruments, drums, bass, etc.
from the rest of the sounds.
To achieve this we use [MDX Net](https://arxiv.org/abs/2111.12203) neural networks (models).
**AudioSeparation** currently supports [39 models](https://huggingface.co/set-soft/audio_separation)
collected by the [UVR5](https://github.com/Anjok07/ultimatevocalremovergui) project.
To achieve this we use [MDX Net](https://arxiv.org/abs/2111.12203) and
[Demucs](https://github.com/facebookresearch/demucs) neural networks (models).
**AudioSeparation** currently supports [46 models](https://huggingface.co/set-soft/audio_separation)
mostly collected by the [UVR5](https://github.com/Anjok07/ultimatevocalremovergui) project.
The models are small (from 21 MB to 65 MB), but really efficient.
Models specialized on different stems are provided.
The MDX models are small (from 21 MB to 65 MB), but really efficient, the Demucs models are bigger
(from 84 MB to 870) MB, but slightly better, and supports 4 stems.
MDX models specialized on different stems are provided.
We support more than one model for each task because some times a model will perform better
for a song and worst for others.
@@ -28,12 +30,19 @@ to keep the secondary vocals along with the instruments.
The objectives for these nodes are:
- Multiple stems (Vocals, Instruments, Drums, Bass, etc.)
- Easy of use
- Clear download (with progress and known destination)
- Support for all possible input audio formats (mono/stereo, any sample rate, any batch size)
- Good quality vs size
- Reduced dependencies
✅ Multiple stems (Vocals, Instruments, Drums, Bass, etc.)
✅ Easy of use
✅ Clear download (with progress and known destination)
✅ Support for all possible input audio formats (mono/stereo, any sample rate, any batch size)
✅ Good quality vs size, or you can choose better quality using Demucs
✅ Reduced dependencies, just my helpers when using ComfyUI
✅ Multiple examples
---
@@ -50,7 +59,10 @@ The objectives for these nodes are:
* [Command Line](#command-line)
* ✨ [Nodes](#-nodes)
* [Vocals using MDX](#vocals-using-mdx)
* [Demucs Audio Separator](#demucs-audio-separator)
* 🖼️ [Examples](#️-examples)
* 📝 [Usage Notes](#-usage-notes)
* 📜 [Project History](#-project-history)
* ⚖️ [License](#️-license)
* 🙏 [Attributions](#-attributions)
@@ -68,10 +80,11 @@ or just do it manually:
cd ComfyUI/custom_nodes/
git clone https://github.com/set-soft/AudioSeparation
```
2. Restart ComfyUI.
2. Install SeCoNoHe: `pip install seconohe`
3. Restart ComfyUI.
The nodes should then appear under the "audio/separation" category in the "Add Node" menu.
You don't need to install extra dependencies.
SeCoNoHe are just a bunch of helpers I created with common functionality I use in my nodes.
### Command Line Tool
@@ -94,13 +107,7 @@ pip install -r requirements.txt
4. Run the scripts like this:
```
python3 tool/demix.py AUDIO_FILE
````
or
```
python tool/demix.py AUDIO_FILE
tool/demix.py AUDIO_FILE
````
You don't need to install it, you could even add a symlink in `/usr/bin`.
@@ -111,7 +118,7 @@ A list of all the available tools can be found [here](tool/README.md).
Models are automatically downloaded.
When using ComfyUI they are downloaded to `ComfyUI/models/audio/MDX`.
When using ComfyUI they are downloaded to `ComfyUI/models/audio/MDX` and `ComfyUI/models/audio/Demucs`.
When using the command line the default is `../models` relative to the script, but you can specify another dir.
@@ -123,7 +130,7 @@ For the command line you can also download the ONNX files from other repos.
## 📦 Dependencies
These nodes just uses `torchaudio` (part of PyTorch), `numpy` for math and `tqdm` for progress bars.
These nodes just uses `torchaudio` (part of PyTorch), `numpy` for math, `safetensors` to load models and `tqdm` for progress bars.
All of them are used by ComfyUI, so you don't need to install any additional dependency on a ComfyUI setup.
The following are optional dependencies:
@@ -141,13 +148,14 @@ You can start using template workflows, go to the ComfyUI *Workflow* menu and th
look for *Audio Separation*
If you want to do it manually you'll find the nodes in the *audio/separation* category.
Or you can use the search menu, double click in the canvas and then type **MDX**:
Or you can use the search menu, double click in the canvas and then type **MDX** (or **Demucs**):
![ComfyUI Search](doc/node_search.png)
Choose a node to extract what you want, i.e. *Vocals*. The complement output for it will
be the instruments, but using a node for *Instrumental* separation you'll usually get a better result
than using the *Complement* output.
than using the *Complement* output. In the case of *Demucs* models you get 4 or 6 stems at a time,
the "UVR Demucs" is an exception, it just supports Vocals and Other.
Then simply connect your audio input to the node (i.e. **LoadAudio** node from Comfy core) and
connect its output to some audio node (i.e. **PreviewAudio** or **SaveAudio** nodes from Comfy core).
@@ -207,6 +215,59 @@ share the same structure, so here is the first:
- **Input Batch Handling:** If `input_sound` is a batch the outputs will be batches. The process is sequential, not parallel.
- **Missing Models:** They are downloaded and stored under `models/audio/MDX` of the ComfyUI installation
And here is the Demucs node:
### Demucs Audio Separator
- **Display Name:** `Demucs Audio Separator`
- **Internal Name:** `AudioSeparateDemucs`
- **Category:** `audio/separation`
- **Description:** Takes one audio input (which can be a batch) separates the vocals, drums and bass from the rest of the sounds.
The node has outputs for guitar and piano, which can be separated by the *Hybrid Transformer 6 sources* model, which is quite
experimental.
- **Inputs:**
- `input_sound` (AUDIO): The audio input. Can be a single audio item or a batch.
- `model` (COMBO): The name of the model to use. Choose one from the list.
- `shifts` (INT): Number of random shifts for equivariant stabilization.
It does extra passes using slightly shifted audio, which can produce better results.
Higher values improve quality but are slower. 0 disables it.
- `overlap` (FLOAT): Amount of overlap between audio chunks.
This is expressed as a portion of the total chunk, i.e. 0.25 is 25%.
Higher values can reduce stitching artifacts but are slower.
- `custom_segment` (BOOLEAN): Enable to override the model's default segment length.
Disabling uses the recommended length from the model file.
Useful for HDemucs and Demucs models, not much for HTDemucs.
- segment (INT): Length of audio chunks to process at a time (in seconds).
Higher values need more VRAM but can improve quality.
- `taget_device` (COMBO): The device where we will run the neural network.
- **Output:**
- `Vocals` (AUDIO): The separated vocals
- `Drums` (AUDIO): The separated drums. Not for "UVR" version.
- `Bass` (AUDIO): The separated bass. Not for "UVR" version.
- `Other` (AUDIO): The separated stuff that doesn't fit in the other outputs
- `Guitar` (AUDIO): The separated guitar, only for *Hybrid Transformer 6 sources* model, which is quite
- `Piano` (AUDIO): The separated piano, only for *Hybrid Transformer 6 sources* model, which is quite
- **Behavior Details:**
- **Sample Rate:** The sample rate of the input is adjusted to 44.1 kHz
- **Channels:** Mono audios are converted to fake stereo (left == right)
- **Input Batch Handling:** If `input_sound` is a batch the outputs will be batches.
- **Missing Models:** They are downloaded and stored under `models/audio/Demucs` of the ComfyUI installation
- **Models:** Note that most models are a *bag of models*, this is four models working together.
## 🖼️ Examples
Once installed the examples are available in the ComfyUI workflow templates, in the *audio-separation* section.
Note that we have two versions, the regular and the *quick* version. The *quick* version is ideal for quick tests, the input files
are downloaded and, in most cases, only 10 seconds of audio are processed.
- [00_Vocals_simple.json](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/00_Vocals_simple.json) [Quick](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/00_Vocals_quick.json): Example to get vocals using MDX
- [01_Vocals_Drums_Bass.json](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/01_Vocals_Drums_Bass.json) [Quick](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/01_Vocals_Drums_Bass_quick.json): Example to get vocals, drums, bass and others using MDX
- [02_Batch.json](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/02_Batch.json) [Quick](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/02_Batch_quick.json): Shows how to apply MDX demix to a batch of audios, using *Audio Batch* nodes.
- [03_Instrumental_keep.json](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/03_Instrumental_keep.json) [Quick](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/03_Instrumental_keep_quick.json): Shows how to extract vocals maintaining the same number of channels and sample rate, using *Audio Batch* nodes.
- [04_Demucs.json](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/04_Demucs.json) [Quick](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/04_Demucs_quick.json): Separates vocals, drums, bass and others using Demucs, better quality.
- [05_Demix_and_Remix.json](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/05_Demix_and_Remix.json) [Quick](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/05_Demix_and_Remix_quick.json): Example to separate vocals and others to the left channel and drums and bass to the right channel, using *Audio Batch* nodes.
## 📝 Usage Notes
@@ -217,6 +278,17 @@ share the same structure, so here is the first:
You can control log verbosity through ComfyUI's startup arguments (e.g., `--preview-method auto --verbose DEBUG` for more detailed ComfyUI logs
which might also affect custom node loggers if they are configured to inherit levels). The logger name used is "AudioSeparation".
You can force debugging level for these nodes defining the `AUDIOSEPARATION_NODES_DEBUG` environment variable to `1`.
- **Models format:** We use safetensors because this format is safer than PyTorch files (.pth, .th, etc.) and doesn't need an extra runtime (like ONNX does)
- **No quantized Demucs:** These models just save download time, but pulls extra dependency (diffq), they are just lower quality versions of their non-quantized counterparts.
## 📜 Project History
- 1.0.0 2025-07-02: Initial release. MDX-Net models support
- 1.1.0 2025-07-11: Demucs models support
- 1.1.1 2025-07-21: More examples. One more Demucs model
## ⚖️ License
@@ -237,6 +309,7 @@ ______
- Models collected by the [UVR5 project](https://github.com/Anjok07/ultimatevocalremovergui) and
found in the [UVR Resources](https://huggingface.co/Politrees/UVR_resources) by
[Artyom Bebroy](https://github.com/Politrees)
- Demucs models are from Meta Platforms, Inc. Except for the "UVR" version
- The logo image was created using text generated using [Text Studio](https://www.textstudio.com/) and
resources from [Vecteezy](https://www.vecteezy.com/) by:
- [Titima Ongkantong](https://www.vecteezy.com/members/titima157)
+6 -26
View File
@@ -3,33 +3,13 @@
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
import inspect
import logging
from .source.utils.misc import NODES_NAME
from . import nodes # noqa: E402
init_logger = logging.getLogger(f"{NODES_NAME}.__init__")
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
# This is our first import so we initialize SeCoNoHe
from .src.nodes import nodes, main_logger, __version__
from seconohe.register_nodes import register_nodes
from seconohe import JS_PATH
def register_nodes(module):
suffix = " " + module.SUFFIX if hasattr(module, "SUFFIX") else ""
if suffix:
suffix = " " + suffix
for name, obj in inspect.getmembers(module):
if not inspect.isclass(obj) or not hasattr(obj, "INPUT_TYPES"):
continue
assert hasattr(obj, "UNIQUE_NAME"), f"No name for {obj.__name__}"
NODE_CLASS_MAPPINGS[obj.UNIQUE_NAME] = obj
NODE_DISPLAY_NAME_MAPPINGS[obj.UNIQUE_NAME] = obj.DISPLAY_NAME + suffix
register_nodes(nodes)
init_logger.info(f"Registering {len(NODE_CLASS_MAPPINGS)} node(s).")
init_logger.debug(f"{list(NODE_DISPLAY_NAME_MAPPINGS.values())}")
WEB_DIRECTORY = "./js"
NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS = register_nodes(main_logger, [nodes], version=__version__)
WEB_DIRECTORY = JS_PATH
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
Binary file not shown.

After

Width:  |  Height:  |  Size: 49 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 18 KiB

+28
View File
@@ -0,0 +1,28 @@
# Audio Separation assets
## Icon
![Icon 1024x1024](AudioSeparation_1024.jpg)
prompt:
```
Icon design, a single, vibrant white soundwave enters a sleek, crystalline triangular prism from the left.
The prism refracts the soundwave, causing it to split into four distinct, brilliantly colored waveforms exiting to the right.
One waveform is glowing red for vocals, one is electric blue for drums, one is deep green for bass, and one is bright yellow for other instruments.
Minimalist, vector art, graphic illustration, glowing neon lines, on a clean dark background. Visually compelling UI icon, 400x400.
```
Model: HiDream I1
Version used:
![Icon 400x400](AudioSeparation_400.jpg)
## Banner
![Banner 21:9](audioseparation_logo_21_9.jpg)
Image composition
Binary file not shown.

After

Width:  |  Height:  |  Size: 69 KiB

+11
View File
@@ -0,0 +1,11 @@
# Internals
## Call sequence
- The node creates a demixer class with: get_demixer(model_data, device, models_dir)
- This function creates a Demixer object for the correct demix type
- The DemixerGeneric.__init__() calls load_model
- load_model:
- Calls get_model to get a proper model object without weights
- Calls a loader for the container (ONNX, safetensors)
- load_safetensors loads the weights
File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
audioseparation_logo.jpg
File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
audioseparation_logo.jpg
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
audioseparation_logo.jpg
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
audioseparation_logo.jpg
File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
audioseparation_logo.jpg
File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
audioseparation_logo.jpg
File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
audioseparation_logo.jpg
File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
audioseparation_logo.jpg
File diff suppressed because one or more lines are too long
-57
View File
@@ -1,57 +0,0 @@
// Copyright (c) 2025 Salvador E. Tropea
// Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
// License: GPLv3
// Project: ComfyUI-AudioSeparation
// This script adds an event named "set-audioseparation-node"
// It can currently just modify a widget value for the current node
import { app } from "/scripts/app.js";
// Register a new extension
app.registerExtension({
name: "SET.AudioSeparation.NodeAdjust", // Unique name
// The setup function is executed when the extension is loaded
setup() {
// Add a listener for our custom event
app.api.addEventListener("set-audioseparation-node", (event) => {
// The data from Python is in event.detail
const { action, arg1, arg2 } = event.detail;
// Find the node that is currently being executed
const node = app.graph.getNodeById(app.runningNodeId);
if (!node) {
console.warn(`[SET.AudioSeparation] Could not find running node with ID: ${app.runningNodeId}`);
return;
}
// --- ACTION EXECUTED HERE ---
switch (action) {
case 'change_widget':
// arg1 = widget name (e.g., "model")
// arg2 = new value (e.g., "💾 My Awesome Model")
const widget = node.widgets.find(w => w.name === arg1);
if (widget) {
// This is the key part for combo boxes (dropdowns)
// If the new value isn't in the list of options, add it first.
if (!widget.options.values.includes(arg2)) {
widget.options.values.push(arg2);
}
// Set the widget value
widget.setValue(arg2, node, app.canvas);
} else {
console.error(`[SET.AudioSeparation] Widget '${arg1}' not found on node ${node.id}`);
}
break;
// Other actions here in the future
// case 'disable_widget':
// ...
// break;
}
});
},
});
-32
View File
@@ -1,32 +0,0 @@
// Copyright (c) 2025 Salvador E. Tropea
// Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
// License: GPLv3
// Project: ComfyUI-AudioSeparation
// This script adds an event named "set-audioseparation-toast"
// Used to notify the user in the GUI using the Toast API
import { app } from "/scripts/app.js";
// Register a new extension
app.registerExtension({
name: "SET.AudioSeparation.ToastHandler", // Unique name
// The setup function is executed when the extension is loaded
setup() {
// Add a listener for our custom event
app.api.addEventListener("set-audioseparation-toast", (event) => {
// The data from Python is in event.detail
const { message, summary, severity } = event.detail;
// Use the ComfyUI toast API to show the message
// app.ui.toast.addMessage is the modern way to do this
app.extensionManager.toast.add({
severity: severity,
summary: summary,
detail: message,
life: 6000
});
});
},
});
+311
View File
@@ -5,6 +5,20 @@
"063aadd735d58150722926dcbf5852a9": {
"config_yaml": "model_2_stem_061321.yaml"
},
"09c249b06a99f569ae2052564659de88": {
"desc": "Hybrid Transformer Demucs fine-tuned",
"download": "Main/Demucs",
"file_t": "safetensors",
"model_t": "Demucs",
"name": "htdemucs_ft.safetensors",
"params": "167937824",
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other"
]
},
"0ddfc0eb5792638ad5dc27850236c246": {
"channels": 32,
"compensate": 1.035,
@@ -90,6 +104,20 @@
"1e6165b601539f38d0a9330f3facffeb": {
"config_yaml": "model_2_stem_061321.yaml"
},
"1f1dabb6daf9e306d3dfadeb89f5370f": {
"desc": "Hybrid Demucs v3, retrained",
"download": "Main/Demucs",
"file_t": "safetensors",
"model_t": "Demucs",
"name": "hdemucs_mmi.safetensors",
"params": "83637832",
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other"
]
},
"203f2a3955221b64df85a41af87cf8f0": {
"compensate": 1.035,
"mdx_dim_f_set": 3072,
@@ -244,6 +272,19 @@
"primary_stem": "Vocals",
"stages": 5
},
"3b90a44a87135fb7b0ac40ed8ee3e16a": {
"desc": "Hybrid Demucs MDX 2nd 2021B",
"download": "FBDemucs/mdx_final",
"file_t": "pytorch",
"model_t": "Demucs",
"name": "mdx_extra.yaml",
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other"
]
},
"3bff56e6709357854e71cb2e7802733a": {
"config_yaml": "config_dnr_bandit_bsrnn_multi_mus64.yaml",
"is_karaoke": false,
@@ -353,6 +394,19 @@
"primary_stem": "Vocals",
"stages": 5
},
"4bbb0edee1a61e7b0f0cd547f3c71af5": {
"desc": "Hybrid Demucs v3, retrained",
"download": "FBDemucs/mdx_final",
"file_t": "pytorch",
"model_t": "Demucs",
"name": "hdemucs_mmi.yaml",
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other"
]
},
"4c0736aa53894dfc6e10f8d178bb8690": {
"channels": 48,
"compensate": 1.035,
@@ -383,6 +437,19 @@
"primary_stem": "Vocals",
"stages": 5
},
"4cca48cc43b93a35fd1e32f70371d5ae": {
"desc": "Hybrid Demucs MDX 1st 2021A",
"download": "FBDemucs/mdx_final",
"file_t": "pytorch",
"model_t": "Demucs",
"name": "mdx.yaml",
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other"
]
},
"4ffce4487f6372bafc2748b2dcdb893f": {
"channels": 32,
"compensate": 1.035,
@@ -476,6 +543,21 @@
"primary_stem": "Vocals",
"stages": 5
},
"5ee7b421b6228a90e63855b108ad830b": {
"desc": "Hybrid Demucs MusDB High Train",
"download": "Main/Demucs",
"file_t": "safetensors",
"model_t": "Demucs",
"name": "hdemucs_high_trained.safetensors",
"params": "83639368",
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other"
],
"use_demucs_pt_process": true
},
"5f6483271e1efb9bfb59e4a3e6d4d098": {
"channels": 32,
"compensate": 1.035,
@@ -524,6 +606,24 @@
"primary_stem": "Other",
"stages": 5
},
"65dd0f678ed827e4e434d97255639c38": {
"desc": "Hybrid Demucs MDX A Rep. TO",
"download": "Main/Demucs",
"file_t": "safetensors",
"is_bag_of_models": "true",
"model_t": "Demucs",
"name": "repro_mdx_a_time_only.yaml",
"params": "267506128",
"parent": "e2bccab7f9797842079390749eee8178",
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other"
],
"segment": "44",
"signatures": "[\"9a6b4851\", \"9a6b4851\", \"1ef250f1\", \"1ef250f1\"]"
},
"6703e39f36f18aa7855ee1047765621d": {
"channels": 32,
"compensate": 1.035,
@@ -575,6 +675,35 @@
"is_roformer": true,
"model_type": "SCNet"
},
"6f82220b9fc8445331d3a246f4ab509a": {
"desc": "Hybrid Demucs MusDB High Train",
"download": "Torch/Demucs",
"file_t": "pytorch",
"model_t": "Demucs",
"name": "hdemucs_high_trained.yaml",
"params": 83639368,
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other"
],
"signatures": "[\"hdemucs_high_trained\"]"
},
"7205f4e53ec3226681ecd0c38e2fa747": {
"desc": "Hybrid Transformer Demucs",
"download": "Main/Demucs",
"file_t": "safetensors",
"model_t": "Demucs",
"name": "htdemucs.safetensors",
"params": "41984456",
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other"
]
},
"73492b58195c3b52d34590d5474452f6": {
"channels": 48,
"compensate": 1.043,
@@ -611,6 +740,19 @@
"is_roformer": true,
"model_type": "SCNet"
},
"80b3c6fe5b8d60a191d7e4ee9029cd3b": {
"desc": "Hybrid Transformer Demucs fine-tuned",
"download": "FBDemucs/hybrid_transformer",
"file_t": "pytorch",
"model_t": "Demucs",
"name": "htdemucs_ft.yaml",
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other"
]
},
"8318a54fe1278ddcf78aad32145c0a6f": {
"config_yaml": "deverb_bs_roformer_8_256dim_8depth.yaml",
"is_karaoke": false,
@@ -679,6 +821,24 @@
"mdx_n_fft_scale_set": 5120,
"primary_stem": "Instrumental"
},
"8abca393bc679ca0336a0ba8ca0e732d": {
"desc": "Hybrid Demucs MDX A Rep. HO",
"download": "Main/Demucs",
"file_t": "safetensors",
"is_bag_of_models": "true",
"model_t": "Demucs",
"name": "repro_mdx_a_hybrid_only.yaml",
"params": "167275664",
"parent": "e2bccab7f9797842079390749eee8178",
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other"
],
"segment": "44",
"signatures": "[\"fa0cb7f9\", \"902315c2\", \"fa0cb7f9\", \"902315c2\"]"
},
"8b15e58e9f33ee39346be7f636ad1d63": {
"channels": 32,
"compensate": 1.03,
@@ -729,6 +889,16 @@
"99b6ceaae542265a3b6d657bf9fde79f": {
"config_yaml": "model_2_stem_full_band_8k.yaml"
},
"9a0be129a16a22799d5587650cf93c0b": {
"desc": "UVR Hybrid Demucs Bag",
"file_t": "pytorch",
"model_t": "Demucs",
"name": "UVR_Demucs_Model_Bag.yaml",
"primary_stem": [
"Vocals",
"Other"
]
},
"9b806a49eb1d25c0dfa14a39ed0c51b8": {
"channels": 32,
"compensate": 1.035,
@@ -1031,12 +1201,42 @@
"primary_stem": "Bass",
"stages": 5
},
"c50d9f98739139dbc9fd2477931d686e": {
"desc": "UVR Hybrid Demucs 2",
"download": "Main/Demucs",
"file_t": "safetensors",
"model_t": "Demucs",
"name": "UVR_Demucs_Model_2.yaml",
"params": "83633212",
"parent": "e4e0c3604695bafc2227555268e37942",
"primary_stem": [
"Vocals",
"Other"
],
"segment": "44",
"signatures": "[\"ebf34a2d\"]"
},
"c7500d7fdb1c0fc24b14b698515462d2": {
"config_yaml": "config_mdx23c_similarity.yaml",
"is_karaoke": false,
"is_roformer": false,
"model_type": "MDX23C"
},
"c8ecb80b71516f0e450a50a23b919f6d": {
"desc": "Hybrid Transformer Demucs 6 sources",
"download": "FBDemucs/hybrid_transformer",
"file_t": "pytorch",
"model_t": "Demucs",
"name": "htdemucs_6s.yaml",
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other",
"Piano",
"Guitar"
]
},
"c93da97c0df6c6f20aa2504bfa2300a0": {
"channels": 48,
"compensate": 1.075,
@@ -1100,6 +1300,21 @@
"primary_stem": "Instrumental",
"stages": 5
},
"ccdd9985f1f7cbe62bf1e1a3abfcbd3c": {
"desc": "UVR Hybrid Demucs 1",
"download": "Main/Demucs",
"file_t": "safetensors",
"model_t": "Demucs",
"name": "UVR_Demucs_Model_1.yaml",
"params": "83633212",
"parent": "e4e0c3604695bafc2227555268e37942",
"primary_stem": [
"Vocals",
"Other"
],
"segment": "44",
"signatures": "[\"ebf34a2db\"]"
},
"cd5b2989ad863f116c855db1dfe24e39": {
"channels": 48,
"compensate": 1.035,
@@ -1219,6 +1434,50 @@
"primary_stem": "Drums",
"stages": 5
},
"dd72197c4f8987a7a033713fa15f0577": {
"desc": "Hybrid Transformer Demucs 6 sources",
"download": "Main/Demucs",
"file_t": "safetensors",
"model_t": "Demucs",
"name": "htdemucs_6s.safetensors",
"params": "27414996",
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other",
"Piano",
"Guitar"
]
},
"e2bccab7f9797842079390749eee8178": {
"desc": "Hybrid Demucs MDX A Reprod.",
"download": "Main/Demucs",
"file_t": "safetensors",
"model_t": "Demucs",
"name": "repro_mdx_a.safetensors",
"params": "434781792",
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other"
]
},
"e2e55d3b30d5f6b345e3e81340542f56": {
"desc": "Hybrid Demucs MDX 2nd 2021B",
"download": "Main/Demucs",
"file_t": "safetensors",
"model_t": "Demucs",
"name": "mdx_extra.safetensors",
"params": "334543632",
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other"
]
},
"e3de6d861635ab9c1d766149edd680d6": {
"config_yaml": "model1.yaml"
},
@@ -1226,6 +1485,18 @@
"config_yaml": "config_vocals_mel_band_roformer_kim.yaml",
"is_roformer": true
},
"e4e0c3604695bafc2227555268e37942": {
"desc": "UVR Hybrid Demucs Bag",
"download": "Main/Demucs",
"file_t": "safetensors",
"model_t": "Demucs",
"name": "UVR_Demucs_Model_Bag.safetensors",
"params": "167266424",
"primary_stem": [
"Vocals",
"Other"
]
},
"e5572e58abf111f80d8241d2e44e7fa4": {
"channels": 48,
"compensate": 1.028,
@@ -1259,6 +1530,20 @@
"e7a25f8764f25a52c1b96c4946e66ba2": {
"config_yaml": "sndfx.yaml"
},
"e94afebb192786ecb9c2111c4eac0835": {
"desc": "Hybrid Demucs MDX 1st 2021A",
"download": "Main/Demucs",
"file_t": "safetensors",
"model_t": "Demucs",
"name": "mdx.safetensors",
"params": "345460200",
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other"
]
},
"e9b82ec90ee56c507a3a982f1555714c": {
"config_yaml": "model_2_stem_full_band_2.yaml"
},
@@ -1269,6 +1554,19 @@
"mdx_n_fft_scale_set": 6144,
"primary_stem": "Instrumental"
},
"ed35ab1c2a2ca529c140636927051b71": {
"desc": "Hybrid Demucs MDX A Reprod.",
"download": "FBDemucs/mdx_final",
"file_t": "pytorch",
"model_t": "Demucs",
"name": "repro_mdx_a.yaml",
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other"
]
},
"eedf36b526be391ad07b04ad4493854d": {
"channels": 48,
"compensate": 1.01,
@@ -1344,6 +1642,19 @@
"primary_stem": "Instrumental",
"stages": 5
},
"f9416d8432ab7ee47c4676bf0c1f8aa7": {
"desc": "Hybrid Transformer Demucs",
"download": "FBDemucs/hybrid_transformer",
"file_t": "pytorch",
"model_t": "Demucs",
"name": "htdemucs.yaml",
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other"
]
},
"fbbf86f7d863a6df7fb5e19b939a9f58": {
"channels": 32,
"compensate": 1.035,
-143
View File
@@ -1,143 +0,0 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
import os
import torch
from typing import Dict
# ComfyUI imports
import folder_paths # ComfyUI's way to access model paths
# Local imports
from .source.utils.logger import main_logger
from .source.utils.load_audio import audio_get_channels, force_stereo, force_sample_rate
from .source.utils.torch import get_torch_device_options
from .source.utils.comfy_node_action import send_node_action
from .source.inference.demixer import get_demixer
from .source.db.models_db import ModelsDB
DEF_MODEL = 'Kim_Vocal_2.safetensors'
DEF_ENTRY = 'Default'
MODELS_DIR = os.path.join(folder_paths.models_dir, "audio", "MDX")
models_db = ModelsDB(MODELS_DIR)
logger = main_logger
class AudioSeparateVocals:
PRIMARY_STEM = 'Vocals'
MODEL_T = 'MDX'
FILE_T = 'safetensors'
DEFAULT_MODEL = "Kim_Vocal_2.safetensors"
@classmethod
def _get_available_audio_models(cls):
global models_db
# Refresh the database
models_db.refresh()
# Filter the models this node can handle
cls.models_filtered = models_db.get_filtered(primary_stem=cls.PRIMARY_STEM, model_t=cls.MODEL_T, file_t=cls.FILE_T,
default=cls.DEFAULT_MODEL, repeat_dl=True)
# We add any model downloaded and memorized by the GUI
return cls.models_filtered.get_display_names()
@classmethod
def INPUT_TYPES(cls):
device_options, default_device = get_torch_device_options()
return {
"required": {
"input_sound": ("AUDIO",),
"model": (cls._get_available_audio_models(),), # Dropdown for model selection
"segments": ("INT", {
"default": 1, # Default value
"min": 1, # Minimum allowed value
"max": 64, # Maximum allowed value (set a reasonable practical max)
"step": 1, # Step for slider/spinbox
"display": "slider" # How to display: "number" or "slider"
}),
"target_device": (device_options, {
"default": default_device,
"tooltip": "The device (CPU or CUDA) to which the projection layer will be assigned for computation."}),
}
}
RETURN_TYPES = ("AUDIO", "AUDIO",)
RETURN_NAMES = (PRIMARY_STEM, "Complement",)
FUNCTION = "execute"
CATEGORY = "audio/separation"
DESCRIPTION = "Separates vocals using MDX-Net networks"
UNIQUE_NAME = "AudioSeparateVocals"
DISPLAY_NAME = "Vocals using MDX"
def __init__(self):
super().__init__()
self.demixer = None
def execute(self, input_sound: Dict, model: str, segments: int, target_device: str):
# Get information for the selected model
main_logger.info(f"Selected model: {model}")
model_data = self.models_filtered.get_by_display_name(model)
if model_data is None:
raise ValueError("Unknown model selected, please refresh pressing `R` and select another")
model_path = model_data.get('model_path')
# Create or recycle a demixer
device = torch.device(target_device)
if self.demixer is None or self.demixer.d['hash'] != model_data['hash']:
# New demixer
logger.debug("Creating a new demixer object")
# This will load the model, optionally downloading it
self.demixer = get_demixer(model_data, device, MODELS_DIR)
# Handle a change in the icon of the model name
if model_path is None:
# Was downloaded
send_node_action("change_widget", "model", model_data['indicator'] + model_data['filtered_name'])
# Match channels and S/R
waveform = input_sound['waveform']
sample_rate = input_sound['sample_rate']
if audio_get_channels(waveform) == 1 and self.demixer.ch == 2:
waveform = force_stereo(waveform)
if sample_rate != self.demixer.sr:
waveform = force_sample_rate(waveform, sample_rate, self.demixer.sr)
# Demix
wavs = self.demixer(waveform, segments)
return (wavs[0], wavs[1],)
class AudioSeparateInstrumental(AudioSeparateVocals):
PRIMARY_STEM = 'Instrumental'
DEFAULT_MODEL = "Kim_Inst.safetensors"
DESCRIPTION = "Separates instruments using MDX-Net networks"
UNIQUE_NAME = "AudioSeparateInstrumental"
DISPLAY_NAME = "Instrumental using MDX"
RETURN_NAMES = (PRIMARY_STEM, "Complement",)
class AudioSeparateBass(AudioSeparateVocals):
PRIMARY_STEM = 'Bass'
DEFAULT_MODEL = "kuielab_b_bass.safetensors"
DESCRIPTION = "Separates bass using MDX-Net networks"
UNIQUE_NAME = "AudioSeparateBass"
DISPLAY_NAME = "Bass using MDX"
RETURN_NAMES = (PRIMARY_STEM, "Complement",)
class AudioSeparateDrums(AudioSeparateVocals):
PRIMARY_STEM = 'Drums'
DEFAULT_MODEL = "kuielab_b_drums.safetensors"
DESCRIPTION = "Separates drums using MDX-Net networks"
UNIQUE_NAME = "AudioSeparateDrums"
DISPLAY_NAME = "Drums using MDX"
RETURN_NAMES = (PRIMARY_STEM, "Complement",)
class AudioSeparateVarious(AudioSeparateVocals):
PRIMARY_STEM = ["Other", "Reverb"]
DEFAULT_MODEL = "Reverb_HQ_By_FoxJoy.safetensors"
DESCRIPTION = "Misc. separators using MDX-Net networks"
UNIQUE_NAME = "AudioSeparateVarious"
DISPLAY_NAME = "Various using MDX"
RETURN_NAMES = ("Main", "Complement",)
+17 -3
View File
@@ -1,13 +1,27 @@
[project]
name = "audio-separation"
description = "Audio separation (aka demixing) nodes, for Vocals, Instruments, Bass, Drums and Others. Using MDX-Net, no extra dependencies, support for batch and resample."
version = "1.0.0"
description = """
Audio separation (aka demixing) nodes, for Vocals, Instruments, Bass, Drums and Others (experimental Piano and Guitar).
Using MDX-Net and Demucs, no extra dependencies, support for batch and resample.
Choose between High Quality and Speed. All safetensor models (No ONNX, No PyTorch)
"""
# Inconsistent mechanism needed by comfy-cli, no dynamic variables
version = "1.1.3"
# Deprecated mechanism, comfy-cli doesn't support SPDX
license = { file = "LICENSE" }
dependencies = []
# Not really used, ComfyUI-Manager doesn't use it
# dependencies = ["seconohe>=1.0.2"]
# So we do it in the reverse way ...
dynamic = ["dependencies"]
[project.urls]
Repository = "https://github.com/set-soft/AudioSeparation"
[tool.setuptools.dynamic]
dependencies = {file = ["requirements.txt"]}
[tool.comfy]
PublisherId = "set-soft"
DisplayName = "Audio Separation (Demix)"
Icon = "https://raw.githubusercontent.com/set-soft/AudioSeparation/main/assets/AudioSeparation_400.jpg"
Banner = "https://raw.githubusercontent.com/set-soft/AudioSeparation/main/assets/audioseparation_logo_21_9.jpg"
+2
View File
@@ -1,7 +1,9 @@
torch
torchaudio
numpy
safetensors
tqdm
seconohe>=1.0.2
# Optionals:
# requests
# colorama
-104
View File
@@ -1,104 +0,0 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# Wrappers for the model and inference
import logging
import torch
# ComfyUI imports
try:
import comfy.utils
with_comfy = True
except Exception:
with_comfy = False
# Local imports
from .stft import stft_chunk_process, stft_get_chunks
from ..db.load_model import load_model
from ..utils.misc import NODES_NAME
logger = logging.getLogger(f"{NODES_NAME}.demixer")
SAMPLE_RATE = 44100
def show_inference_parameters(d):
logger.debug("Using inference parameters:")
logger.debug(f" Frequency Bins (n_fft/2): {d['mdx_n_fft_scale_set']//2}")
logger.debug(f" Amplitude Compensation: {d['compensate']}")
class DemixerMDX(object):
def __init__(self, d, device, models_dir):
self.d = d
self.model_run = load_model(d, device, models_dir)
self.device = device
show_inference_parameters(d)
self.sr = SAMPLE_RATE
self.ch = 2
def __call__(self, waveform, segments=1):
dim_t = (2 ** self.d['mdx_dim_t_set']) * segments
try:
# --- 1. Normalize input shape to handle both batched and non-batched data ---
if waveform.ndim == 2:
# Input is [C, samples], add a batch dimension to make it [1, C, samples]
logger.debug("Input is not batched. Adding a temporary batch dimension.")
waveform = waveform.unsqueeze(0)
input_was_batched = False
elif waveform.ndim == 3:
# Input is already batched [B, C, samples]
input_was_batched = True
else:
raise ValueError(f"Unsupported waveform shape: {waveform.shape}. Expected 2 or 3 dimensions.")
batch_size = waveform.shape[0]
logger.info("🎛️ Performing demix...")
# Lists to store the separated stems from each item in the batch
list_of_main_stems = []
list_of_complement_stems = []
# ComfyUI progress bar
progress_bar_ui = None
if with_comfy:
chunks = stft_get_chunks(waveform.shape[2], self.d['mdx_n_fft_scale_set'], segment_size=dim_t)
chunks *= batch_size
progress_bar_ui = comfy.utils.ProgressBar(chunks)
# --- 2. Iterate through the batch ---
for i, single_waveform in enumerate(waveform):
# single_waveform has shape [C, samples]
logger.debug(f"Processing item {i+1}/{batch_size}...")
# Process this single waveform
main_wav = stft_chunk_process(single_waveform, self.d, self.model_run, self.device, segment_size=dim_t,
progress_bar_ui=progress_bar_ui)
complement_wav = single_waveform - main_wav
# Add the results to our lists
list_of_main_stems.append(main_wav)
list_of_complement_stems.append(complement_wav)
# --- 3. Stack the results into single batch tensors ---
# torch.stack creates a new dimension (the batch dimension) from a list of tensors
stacked_main_stems = torch.stack(list_of_main_stems, dim=0)
stacked_complement_stems = torch.stack(list_of_complement_stems, dim=0)
# Both will now have shape [B, C, samples]
# --- 4. Denormalize output shape if original input was not batched ---
if not input_was_batched:
logger.debug("Squeezing batch dimension from output to match non-batched input.")
stacked_main_stems = stacked_main_stems.squeeze(0)
stacked_complement_stems = stacked_complement_stems.squeeze(0)
return [{'waveform': stacked_main_stems, 'sample_rate': SAMPLE_RATE, 'stem': self.d['primary_stem']},
{'waveform': stacked_complement_stems, 'sample_rate': SAMPLE_RATE, 'stem': 'Complement'}]
except Exception as e:
logger.error(f"Error during separation: {str(e)}")
raise e
def get_demixer(d, device, models_dir):
# Currently just MDX
return DemixerMDX(d, device, models_dir)
-21
View File
@@ -1,21 +0,0 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# Helper to get a model from the correct class
import logging
from .MDX_Net import MDX_Net
from ..utils.misc import NODES_NAME
logger = logging.getLogger(f"{NODES_NAME}.get_model")
# Currently we have just one type of networks, but this is a clean way to support more, or even test replacements
def get_model(d):
model_t = d['model_t'].lower()
if model_t != "mdx":
msg = f"Unknown model type `{model_t}`"
logger.error(msg)
raise ValueError(msg)
return MDX_Net(dim_f=d['mdx_dim_f_set'], ch=d['channels'], num_stages=d['stages'])
View File
-113
View File
@@ -1,113 +0,0 @@
# Copyright Jonathan Hartley 2013. BSD 3-Clause license, see LICENSE file.
'''
This module generates ANSI character codes to printing colors to terminals.
See: http://en.wikipedia.org/wiki/ANSI_escape_code
'''
import sys
import os
CSI = '\033['
OSC = '\033]'
BEL = '\a'
is_a_tty = sys.stderr.isatty() and os.name == 'posix'
def code_to_chars(code):
return CSI + str(code) + 'm' if is_a_tty else ''
def set_title(title):
return OSC + '2;' + title + BEL
def clear_screen(mode=2):
return CSI + str(mode) + 'J'
def clear_line(mode=2):
return CSI + str(mode) + 'K'
class AnsiCodes(object):
def __init__(self):
# the subclasses declare class attributes which are numbers.
# Upon instantiation we define instance attributes, which are the same
# as the class attributes but wrapped with the ANSI escape sequence
for name in dir(self):
if not name.startswith('_'):
value = getattr(self, name)
setattr(self, name, code_to_chars(value))
class AnsiCursor(object):
def UP(self, n=1):
return CSI + str(n) + 'A'
def DOWN(self, n=1):
return CSI + str(n) + 'B'
def FORWARD(self, n=1):
return CSI + str(n) + 'C'
def BACK(self, n=1):
return CSI + str(n) + 'D'
def POS(self, x=1, y=1):
return CSI + str(y) + ';' + str(x) + 'H'
class AnsiFore(AnsiCodes):
BLACK = 30
RED = 31
GREEN = 32
YELLOW = 33
BLUE = 34
MAGENTA = 35
CYAN = 36
WHITE = 37
RESET = 39
# These are fairly well supported, but not part of the standard.
LIGHTBLACK_EX = 90
LIGHTRED_EX = 91
LIGHTGREEN_EX = 92
LIGHTYELLOW_EX = 93
LIGHTBLUE_EX = 94
LIGHTMAGENTA_EX = 95
LIGHTCYAN_EX = 96
LIGHTWHITE_EX = 97
class AnsiBack(AnsiCodes):
BLACK = 40
RED = 41
GREEN = 42
YELLOW = 43
BLUE = 44
MAGENTA = 45
CYAN = 46
WHITE = 47
RESET = 49
# These are fairly well supported, but not part of the standard.
LIGHTBLACK_EX = 100
LIGHTRED_EX = 101
LIGHTGREEN_EX = 102
LIGHTYELLOW_EX = 103
LIGHTBLUE_EX = 104
LIGHTMAGENTA_EX = 105
LIGHTCYAN_EX = 106
LIGHTWHITE_EX = 107
class AnsiStyle(AnsiCodes):
BRIGHT = 1
DIM = 2
NORMAL = 22
RESET_ALL = 0
Fore = AnsiFore()
Back = AnsiBack()
Style = AnsiStyle()
Cursor = AnsiCursor()
-44
View File
@@ -1,44 +0,0 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# ComfyUI Node actions
import logging
# ComfyUI imports
try:
from server import PromptServer
with_comfy = True
except Exception:
with_comfy = False
# Local imports
from .misc import NODES_NAME
logger = logging.getLogger(f"{NODES_NAME}.comfy_node_action")
def send_node_action(action: str, arg1: str = None, arg2: str = None, sid: str = None):
"""
Sends a node action event to the ComfyUI client.
Args:
action (str): Action to be performed.
arg1 (str): First argument
arg2 (str): Second argument
sid (str, optional): The session ID of the client to send to.
If None, broadcasts to all clients. Defaults to None.
"""
if not with_comfy:
return
try:
PromptServer.instance.send_sync(
"set-audioseparation-node", # This is our custom event name
{
'action': action,
'arg1': arg1,
'arg2': arg2
},
sid
)
except Exception as e:
logger.error(f"when trying to use ComfyUI PromptServer: {e}")
-46
View File
@@ -1,46 +0,0 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# ComfyUI Toast API messages
# Original code from Gemini 2.5 Pro, which was really outdated
# Took ideas from Easy Use nodes and looking at ComfyUI code
import logging
# ComfyUI imports
try:
from server import PromptServer
with_comfy = True
except Exception:
with_comfy = False
# Local imports
from .misc import NODES_NAME
logger = logging.getLogger(f"{NODES_NAME}.comfy_notification")
def send_toast_notification(message: str, summary: str = "Warning", severity: str = "warn", sid: str = None):
"""
Sends a toast notification event to the ComfyUI client.
Args:
message (str): The message content of the toast.
severity (str): The type of toast. Can be 'success' | 'info' | 'warn' | 'error' | 'secondary' | 'contrast'
summary (str): Short explanation
sid (str, optional): The session ID of the client to send to.
If None, broadcasts to all clients. Defaults to None.
"""
if not with_comfy:
return
try:
PromptServer.instance.send_sync(
"set-audioseparation-toast", # This is our custom event name
{
'message': message,
'summary': summary,
'severity': severity
},
sid
)
except Exception as e:
logger.error(f"when trying to use ComfyUI PromptServer: {e}")
-206
View File
@@ -1,206 +0,0 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# Model downloader w/TQDM and ComfyUI progress
# Original code from Gemini 2.5 Pro
import logging
import os
# Requests is better than the core Python urllib, and is a really common package
# But we don't really need it. Lets make it optional:
try:
import requests
with_requests = True
except Exception:
with_requests = False
import urllib
from tqdm import tqdm
# ComfyUI imports
try:
import comfy.utils
with_comfy = True
except Exception:
with_comfy = False
# Local imports
from .misc import NODES_NAME
logger = logging.getLogger(f"{NODES_NAME}.downloader")
def download_model_requests(url: str, save_dir: str, file_name: str):
"""
Downloads a file from a URL with progress bars for both console and ComfyUI.
Args:
url (str): The direct download URL for the file.
save_dir (str): The directory where the file will be saved.
file_name (str): The name of the file to be saved on disk.
"""
full_path = os.path.join(save_dir, file_name)
# Ensure the save directory exists
os.makedirs(save_dir, exist_ok=True)
try:
# Use a streaming request to handle large files and get content length
with requests.get(url, stream=True, timeout=10) as r:
r.raise_for_status() # Raise an exception for bad status codes (4xx or 5xx)
# Get total file size from headers
total_size_in_bytes = int(r.headers.get('content-length', 0))
block_size = 1024 # 1 Kibibyte
# --- Setup Progress Bars ---
# Console progress bar using tqdm
progress_bar_console = tqdm(
total=total_size_in_bytes,
unit='iB',
unit_scale=True,
desc=f"Downloading {file_name}"
)
# ComfyUI progress bar
progress_bar_ui = comfy.utils.ProgressBar(total_size_in_bytes) if with_comfy else None
# --- Download Loop ---
downloaded_size = 0
with open(full_path, 'wb') as f:
for chunk in r.iter_content(chunk_size=block_size):
if chunk: # filter out keep-alive new chunks
chunk_size = len(chunk)
# Update console progress bar
progress_bar_console.update(chunk_size)
# Update ComfyUI progress bar
downloaded_size += chunk_size
if progress_bar_ui:
progress_bar_ui.update(chunk_size) # ProgressBar takes absolute value, but update is incremental
# Write chunk to file
f.write(chunk)
# --- Cleanup ---
progress_bar_console.close()
# Final check to see if download was complete
if total_size_in_bytes != 0 and progress_bar_console.n != total_size_in_bytes:
logger.error("Download failed: Size mismatch.")
# Optional: remove partial file
# os.remove(full_path)
raise IOError(f"Download failed for {file_name}. Expected {total_size_in_bytes} but got "
f"{progress_bar_console.n}")
return full_path
except requests.exceptions.RequestException as e:
logger.error(f"Network error while downloading {file_name}: {e}")
# Clean up partial file if it exists
if os.path.exists(full_path):
try:
os.remove(full_path)
except OSError:
pass
raise
except Exception as e:
logger.error(f"An error occurred during download: {e}")
if os.path.exists(full_path):
try:
os.remove(full_path)
except OSError:
pass
raise
# A simple version implemented using the Python urllib
class Downloader:
def __init__(self, model_path, model_name):
self.model_path = model_path
self.model_name = model_name
self.model_full_name = os.path.join(self.model_path, self.model_name)
# Ensure the directory for the model_path exists before __init__ if used elsewhere
# or create it at the start of download_model
# A TQDM helper class for urlretrieve reporthook
# This is a common pattern for this use case.
class TqdmUpTo(tqdm):
"""
Provides `update_to(block_num, block_size, total_size)`
and updates the TQDM bar.
"""
def __init__(self, unit, unit_scale, unit_divisor, miniters, desc):
super().__init__(unit=unit, unit_scale=unit_scale, unit_divisor=unit_divisor, miniters=miniters, desc=desc)
self.ui_bar = None
self.total = None
def update_to(self, block_num=1, block_size=1, total_size=None):
"""
block_num : int, optional
Number of blocks transferred so far [default: 1].
block_size : int, optional
Size of each block (in tqdm units) [default: 1].
total_size : int, optional
Total size (in tqdm units). If [default: None] remains unchanged.
"""
if total_size is not None and self.total is None:
self.total = total_size
# ComfyUI progress bar
if self.ui_bar is None and with_comfy:
self.ui_bar = comfy.utils.ProgressBar(total_size)
# self.update() will take the *difference* from the last call.
# So we pass the number of new blocks * block_size.
# Since block_num is cumulative, we calculate the new amount.
chunk_size = block_num * block_size - self.n
self.update(chunk_size) # self.n is current progress
if self.ui_bar:
self.ui_bar.update(chunk_size) # ProgressBar takes absolute value, but update is incremental
def download_model(self, url: str):
try:
# Ensure the directory exists
# Use or '.' for current dir if dirname is empty
os.makedirs(self.model_path or '.', exist_ok=True)
# Get filename for tqdm description
filename = self.model_name
# Use TqdmUpTo as a context manager
with self.TqdmUpTo(unit='iB', unit_scale=True, unit_divisor=1024, miniters=1,
desc=f"Downloading {filename}") as t:
# urlretrieve(url, filename=None, reporthook=None, data=None)
# reporthook is called with (block_num, block_size, total_size)
urllib.request.urlretrieve(url, self.model_full_name, reporthook=t.update_to)
# The 'with' statement ensures t.close() is called.
return filename
except urllib.error.URLError as e: # More specific exception for network issues
# Clean up partially downloaded file if an error occurs
if os.path.exists(self.model_full_name):
os.remove(self.model_full_name)
raise Exception(f"An error occurred while downloading the model (URL Error): {e.reason} from {url}")
except Exception as e:
# Clean up partially downloaded file if an error occurs
if os.path.exists(self.model_full_name):
os.remove(self.model_full_name)
raise Exception(f"An unexpected error occurred while downloading the model: {e}")
def download_model_urllib(url: str, save_dir: str, file_name: str):
return Downloader(save_dir, file_name).download_model(url)
def download_model(url: str, save_dir: str, file_name: str, force_urllib: bool = False):
logger.info(f"Downloading model: {file_name}")
logger.info(f"Source URL: {url}")
full_name = os.path.join(save_dir, file_name)
logger.info(f"Destination: {full_name}")
if with_requests and not force_urllib:
download_model_requests(url, save_dir, file_name)
else:
download_model_urllib(url, save_dir, file_name)
logger.info(f"Successfully downloaded {full_name}")
return full_name
-32
View File
@@ -1,32 +0,0 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# Model loader helper
# Original code from Gemini 2.5 Pro
import logging
from safetensors.torch import load_file
from .misc import NODES_NAME
logger = logging.getLogger(f"{NODES_NAME}.load_safetensors")
def load_safetensors(model_path, model_run, device):
logger.info("Loading PyTorch model from .safetensors file...")
# 1. Load the state_dict from the file, EXPLICITLY forcing all tensors onto the CPU.
state_dict = load_file(model_path, device="cpu")
# 2. Load the CPU state_dict into the CPU model. This is now a safe operation.
try:
missing_keys, unexpected_keys = model_run.load_state_dict(state_dict, strict=False)
if missing_keys:
logger.warning(f"Missing keys in state_dict for model_run: {missing_keys}")
if unexpected_keys:
logger.warning(f"Unexpected keys in state_dict for model_run: {unexpected_keys}")
if not missing_keys and not unexpected_keys:
logger.debug("All keys matched successfully.")
except RuntimeError as e:
logger.error(f"RuntimeError during model_run.load_state_dict: {e}")
logger.error("This might indicate a mismatch between saved weights and model architecture.")
raise
return model_run
-105
View File
@@ -1,105 +0,0 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
import os
import sys
import logging
from .misc import NODES_NAME, NODES_DEBUG_VAR
no_colorama = False
try:
from colorama import init as colorama_init, Fore, Back, Style
except ImportError:
no_colorama = True
# If colorama isn't installed use an ANSI basic replacement
if no_colorama:
from .ansi import Fore, Back, Style # noqa: F811
else:
colorama_init()
# Used for tools
standalone_mode = False
white = Fore.WHITE + Style.BRIGHT
yellow = Fore.YELLOW + Style.BRIGHT
red = Fore.RED + Style.BRIGHT
red_alarm = Fore.RED + Back.WHITE + Style.BRIGHT
cyan = Fore.CYAN + Style.BRIGHT
reset = Style.RESET_ALL
# format = "%(asctime)s - %(name)s - %(levelname)s - %(message)s "
# "(%(filename)s:%(lineno)d)"
format = f"[{NODES_NAME} %(levelname)s] %(message)s (%(name)s - %(filename)s:%(lineno)d)"
format_simple = f"[{NODES_NAME}] %(message)s"
FORMATS = {
logging.DEBUG: cyan + format + reset,
logging.INFO: white + format_simple + reset,
logging.WARNING: yellow + format + reset,
logging.ERROR: red + format + reset,
logging.CRITICAL: red_alarm + format + reset
}
format = "[%(levelname)s] %(message)s (%(name)s - %(filename)s:%(lineno)d)"
format_simple = "%(message)s"
if not sys.stdout.isatty():
white = yellow = red = red_alarm = cyan = reset = ""
FORMATS_STANDALONE = {
logging.DEBUG: cyan + format + reset,
logging.INFO: white + format_simple + reset,
logging.WARNING: yellow + format + reset,
logging.ERROR: red + format + reset,
logging.CRITICAL: red_alarm + format + reset
}
class CustomFormatter(logging.Formatter):
"""Logging Formatter to add colors"""
def __init__(self):
super(logging.Formatter, self).__init__()
def format(self, record):
formats = FORMATS_STANDALONE if standalone_mode else FORMATS
log_fmt = formats.get(record.levelno)
formatter = logging.Formatter(log_fmt)
return formatter.format(record)
# Create a new logger
logger = logging.getLogger(NODES_NAME)
logger.propagate = False
# Add handler if we don't have one.
if not logger.handlers:
handler = logging.StreamHandler(sys.stdout)
handler.setFormatter(CustomFormatter())
logger.addHandler(handler)
# ######################
# Logger setup
# ######################
# 1. Determine the ComfyUI global log level (influenced by --verbose)
main_logger = logger
comfy_root_logger = logging.getLogger('comfy')
effective_comfy_level = logging.getLogger().getEffectiveLevel()
# 2. Check our custom environment variable for more verbosity
try:
nodes_debug_env = int(os.environ.get(NODES_DEBUG_VAR, "0"))
except ValueError:
nodes_debug_env = 0
# 3. Set node's logger level
if nodes_debug_env:
main_logger.setLevel(logging.DEBUG - (nodes_debug_env - 1))
final_level_str = f"DEBUG (due to {NODES_DEBUG_VAR}={nodes_debug_env})"
else:
main_logger.setLevel(effective_comfy_level)
final_level_str = logging.getLevelName(effective_comfy_level) + " (matching ComfyUI global)"
_initial_setup_logger = logging.getLogger(NODES_NAME + ".setup") # A temporary logger for this message
_initial_setup_logger.debug(f"{NODES_NAME} logger level set to: {final_level_str}")
def logger_set_standalone(args):
verbose = args.verbose
global main_logger
main_logger.setLevel(logging.DEBUG - (verbose - 1) if verbose else logging.INFO)
global standalone_mode
standalone_mode = True
-18
View File
@@ -1,18 +0,0 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
import logging
NODES_NAME = "AudioSeparation"
NODES_DEBUG_VAR = NODES_NAME.upper() + "_NODES_DEBUG"
def debugl(logger, level, msg):
if logger.getEffectiveLevel() <= logging.DEBUG - (level - 1):
logger.debug(msg)
def cli_add_verbose(parser):
parser.add_argument('-v', '--verbose', action='count', default=0,
help="Enable verbose output to see details of the process.")
-111
View File
@@ -1,111 +0,0 @@
# -*- coding: utf-8 -*-
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: CC BY-NC-SA 4.0
# Project: ComfyUI-Float_Optimized
import contextlib # For context manager
import logging
import torch
try:
import comfy.model_management as mm
with_comfy = True
except Exception:
with_comfy = False
from .misc import NODES_NAME
logger = logging.getLogger(f"{NODES_NAME}.torch")
def get_torch_device_options():
# We always have CPU
default = "cpu"
options = [default]
# Do we have CUDA?
if torch.cuda.is_available():
default = "cuda"
options.append(default)
for i in range(torch.cuda.device_count()):
options.append(f"cuda:{i}") # Specific CUDA devices
# Is this a Mac?
if torch.backends.mps.is_available() and torch.backends.mps.is_built():
options.append("mps")
if default == "cpu":
default = "mps"
return options, default
# ##################################################################################
# # Helper for inference (Target device, offload, eval, no_grad and cuDNN Benchmark)
# ##################################################################################
@contextlib.contextmanager
def model_to_target(model):
"""
Consolidated context manager for model device placement and inference state.
- Moves the model to its designated `model.target_device`.
- Sets `torch.backends.cudnn.benchmark` based on `model.cudnn_benchmark_setting` if available.
- Sets the model to `eval()` mode.
- Wraps the operation in a `torch.no_grad()` context.
- Offloads the model to the CPU (`mm.unet_offload_device()`) afterwards.
"""
if not isinstance(model, torch.nn.Module):
with torch.no_grad():
yield # The code inside the 'with' statement runs here
return
# 1. Determine target device from the model object
try:
target_device = model.target_device
assert isinstance(target_device, torch.device)
except Exception as e:
logger.warning(f"model_to_target: Could not get 'target_device' from model ({e}). "
"Defaulting to model's current device.")
target_device = next(model.parameters()).device
# 2. Get CUDNN benchmark setting from the model object (optional)
# Use hasattr as this is an optional setting that not all models might have.
cudnn_benchmark_enabled = None # Default is to keep the current setting
if hasattr(model, 'cudnn_benchmark_setting'):
cudnn_benchmark_enabled = model.cudnn_benchmark_setting
original_device = next(model.parameters()).device
original_cudnn_benchmark_state = None
is_cuda_target = target_device.type == 'cuda'
try:
# 3. Manage cuDNN benchmark state
if (cudnn_benchmark_enabled is not None and is_cuda_target and hasattr(torch.backends, 'cudnn') and
torch.backends.cudnn.is_available()):
if torch.backends.cudnn.benchmark != cudnn_benchmark_enabled:
original_cudnn_benchmark_state = torch.backends.cudnn.benchmark
torch.backends.cudnn.benchmark = cudnn_benchmark_enabled
logger.debug(f"Temporarily set cuDNN benchmark to {torch.backends.cudnn.benchmark}")
# 4. Move model to target device if not already there
if original_device != target_device:
logger.debug(f"Moving model from `{original_device}` to target device `{target_device}` for inference.")
model.to(target_device)
# 5. Set to eval mode and disable gradients for the operation
model.eval()
with torch.no_grad():
yield # The code inside the 'with' statement runs here
finally:
# 6. Restore original cuDNN benchmark state
if original_cudnn_benchmark_state is not None:
# This check is sufficient because it will only be not None if we set it inside the try block
torch.backends.cudnn.benchmark = original_cudnn_benchmark_state
logger.debug(f"Restored cuDNN benchmark to {original_cudnn_benchmark_state}")
# 7. Offload model back to CPU
if with_comfy:
offload_device = mm.unet_offload_device()
current_device_after_yield = next(model.parameters()).device
if current_device_after_yield != offload_device:
logger.debug(f"Offloading model from `{current_device_after_yield}` to offload device `{offload_device}`.")
model.to(offload_device)
# Clear cache if we were on a CUDA device
if 'cuda' in str(current_device_after_yield):
torch.cuda.empty_cache()
+13
View File
@@ -0,0 +1,13 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPL-3.0
# Project: ComfyUI-AudioSeparation
from seconohe.logger import initialize_logger
__version__ = "1.1.3"
__copyright__ = "Copyright © 2025 Salvador E. Tropea / Instituto Nacional de Tecnología Industrial"
__license__ = "License GPLv3+: GNU GPL version 3 or later <https://gnu.org/licenses/gpl.html>"
__author__ = "Salvador E. Tropea"
NODES_NAME = "AudioSeparation"
main_logger = initialize_logger(NODES_NAME)
@@ -10,10 +10,11 @@
import os
import csv
import logging
from seconohe.logger import debugl
import sys
from .hash import get_hash
from ..utils.misc import NODES_NAME, debugl
from .. import NODES_NAME
# Set up the logger as specified
logger = logging.getLogger(f"{NODES_NAME}.hash_dir")
@@ -80,7 +81,7 @@ def hash_dir(directory_path: str) -> dict:
# 4. Filter by file size.
try:
file_size = os.path.getsize(file_path)
if file_size < MIN_FILE_SIZE_BYTES:
if file_size < MIN_FILE_SIZE_BYTES and not filename.lower().endswith('.yaml'):
logger.debug(f"Skipping small file: '{filename}' ({file_size / 1024**2:.2f}MB)")
continue
except OSError as e:
@@ -132,7 +133,7 @@ def hash_dir(directory_path: str) -> dict:
# ==============================================================================
# Command-Line Tool for Testing and Operation
# python -m source.db.hash_dir
# python -m src.nodes.db.hash_dir
# ==============================================================================
if __name__ == "__main__":
# Local imports to avoid top-level pollution
@@ -9,16 +9,18 @@ import logging
# Local imports
from .models_db import download_model
from ..inference.get_model import get_model
from ..utils.misc import NODES_NAME
from .. import NODES_NAME
logger = logging.getLogger(f"{NODES_NAME}.load_model")
def show_model_parameters(d):
logger.debug("Using model parameters:")
logger.debug(f" Frequency Dimension (dim_f): {d['mdx_dim_f_set']}")
logger.debug(f" Base Channels (ch): {d['channels']}")
logger.debug(f" U-Net Stages: {d['stages']}")
model_t = d['model_t'].lower()
if model_t == 'mdx':
logger.debug("Using model parameters:")
logger.debug(f" Frequency Dimension (dim_f): {d['mdx_dim_f_set']}")
logger.debug(f" Base Channels (ch): {d['channels']}")
logger.debug(f" U-Net Stages: {d['stages']}")
def load_model(d, device, models_dir):
@@ -30,6 +32,16 @@ def load_model(d, device, models_dir):
if model_path is None:
# It means it wasn't on disk
model_path = download_model(d, models_dir)
# Is this a child model?
parent = d.get("parent")
if parent:
# Yes, we need the parent file
if not isinstance(parent, dict):
raise ValueError("Trying to load a broken child model")
model_path = parent.get('model_path')
if model_path is None:
# It means it wasn't on disk
model_path = download_model(parent, models_dir)
# ONNX
if file_t == "onnx":
@@ -9,9 +9,8 @@ import json
import logging
import os
from pathlib import Path
from ..utils.misc import NODES_NAME
from ..utils.downloader import download_model as download_model_basic
from ..utils.comfy_notification import send_toast_notification
from seconohe.downloader import download_file as download_model_basic
from .. import NODES_NAME
from .hash_dir import hash_dir
from .hash import is_hash, get_hash
@@ -22,7 +21,8 @@ known_models_mtime = None
ICON_REMOTE = "⬇️ " # "\u2B07" # ⬇️
ICON_DOWNLOADED = "\U0001F4BE " # 💾 Floppy Disk
KNOWN_SOURCES = {'Politrees/MDXNet': 'https://huggingface.co/Politrees/UVR_resources/resolve/main/models/MDXNet',
'Main/MDX': 'https://huggingface.co/set-soft/audio_separation/resolve/main/MDX'}
'Main/MDX': 'https://huggingface.co/set-soft/audio_separation/resolve/main/MDX',
'Main/Demucs': 'https://huggingface.co/set-soft/audio_separation/resolve/main/Demucs', }
def get_db_filename(provided=None):
@@ -32,7 +32,7 @@ def get_db_filename(provided=None):
script_dir = Path(__file__).resolve().parent
# Build the path to the JSON file: go up one level, then into models/
json_path = script_dir / ".." / ".." / "models" / "uvr_model_data.json"
json_path = script_dir / ".." / ".." / ".." / "models" / "uvr_model_data.json"
try:
return json_path.resolve().relative_to(Path.cwd())
except ValueError:
@@ -68,9 +68,10 @@ def load_known_models(json_path=None):
logger.error("Error: The models database was not found at the expected location.")
logger.error("Please check the directory structure.")
return None
except json.JSONDecodeError:
logger.error(f"Error: The file at '{json_path}' is not a valid JSON file.")
return None
except json.JSONDecodeError as e:
msg = f"Error: The file at '{json_path}' is not a valid JSON file: {e}"
logger.error(msg)
raise ValueError(msg)
except Exception as e:
logger.error(f"An unexpected error occurred: {e}")
return None
@@ -125,6 +126,8 @@ def get_models_full(primary_stem=None, model_t=None, file_t=None, json_path=None
# Allow for multiple values in the filters
if isinstance(primary_stem, str):
primary_stem = {primary_stem}
elif isinstance(primary_stem, list):
primary_stem = set(primary_stem)
if isinstance(model_t, str):
model_t = {model_t}
if isinstance(file_t, str):
@@ -153,8 +156,14 @@ def get_models_full(primary_stem=None, model_t=None, file_t=None, json_path=None
continue
# Check the stem
try:
if primary_stem is not None and d['primary_stem'] not in primary_stem:
continue
if primary_stem is not None:
mps = d['primary_stem']
if isinstance(mps, str):
if mps not in primary_stem:
continue
else: # A list
if not (set(mps) & primary_stem):
continue
except KeyError:
logger.error(f"Missing `primary_stem` for {name}")
continue
@@ -201,12 +210,24 @@ def get_models_full(primary_stem=None, model_t=None, file_t=None, json_path=None
on_disk_as_down.append(ICON_REMOTE + filtered_name)
else:
to_down.append(ICON_REMOTE + filtered_name)
d['hash'] = hash
found[filtered_name] = d
found_hashes[hash] = d
if file_name is not None:
d['model_path'] = downloaded[hash]
found_disk[os.path.realpath(file_name)] = d
# Link the child models
for d in found.values():
parent_hash = d.get("parent")
if not parent_hash or isinstance(parent_hash, dict):
continue
parent_obj = found_hashes.get(parent_hash)
if not parent_obj:
logger.error(f"Inconsistency in database, model `{d['name']}` points to unknown hash `{parent_hash}`")
continue
d["parent"] = parent_obj
return found, found_hashes, found_disk, def_sep + sorted(on_disk) + sorted(to_down) + sorted(on_disk_as_down)
@@ -222,20 +243,25 @@ def get_download_url(data):
except KeyError:
return None
try:
return os.path.join(KNOWN_SOURCES[dn_t], name)
return KNOWN_SOURCES[dn_t] + '/' + name
except KeyError:
logger.error(f"Unknown download source `{dn_t}`")
return None
def cli_add_db(parser, default_json_file=None):
default_json_file = default_json_file or get_db_filename()
parser.add_argument('--json_file', type=str, default=default_json_file,
help="Path to the models database JSON file.")
def cli_add_models_and_db(parser):
# Compute the models dir assuming the script is run from a clone of the repo
default_json_file = get_db_filename()
parser.add_argument('--models_dir', type=str, default=os.path.dirname(default_json_file),
help="Path to the directory containing model files.")
parser.add_argument('--json_file', type=str, default=default_json_file,
help="Path to the models database JSON file.")
cli_add_db(parser, default_json_file=default_json_file)
def download_model(data, models_dir):
@@ -246,15 +272,13 @@ def download_model(data, models_dir):
# Download the file
name = data['name']
send_toast_notification(f"Downloading `{name}`", "Download")
try:
fname = download_model_basic(url, models_dir, name)
fname = download_model_basic(logger, url, models_dir, name)
# Mark it as downloaded
data['model_path'] = fname
if data['indicator']:
data['indicator'] = ICON_DOWNLOADED
# Notify the user
send_toast_notification("Finished downloading", "Download", 'success')
return fname
except Exception as e:
raise ValueError(f"Failed to download {name} from {url}\n{e}")
@@ -305,13 +329,49 @@ class ModelsDB(object):
def __init__(self, models_dir: str, json_path: str = None):
super().__init__()
self.models_dir = models_dir
self.json_path = json_path
self.json_path = get_db_filename(json_path)
self.refresh()
def refresh(self):
self.downloaded = hash_dir(self.models_dir)
self.models = load_known_models(self.json_path)
def remove(self, data):
if data is None:
return
del self.models[data["hash"]]
def add(self, hash, data):
self.models[hash] = data
def save(self):
# Remove run-time information
for k, v in self.models.items():
try:
# Make children just refer to parent's hash, not the actual parent
parent = v["parent"]
if isinstance(parent, dict):
v["parent"] = parent["hash"]
except KeyError:
pass
try:
del v["hash"]
except KeyError:
pass
try:
del v["indicator"]
except KeyError:
pass
try:
del v["filtered_name"]
except KeyError:
pass
try:
del v["model_path"]
except KeyError:
pass
save_known_models(self.models, self.json_path)
def get_filtered(self, primary_stem=None, model_t=None, file_t=None, default=None, repeat_dl=False):
return FilteredModels(primary_stem=primary_stem, model_t=model_t, file_t=file_t, json_path=self.json_path,
downloaded=self.downloaded, default=default, repeat_dl=repeat_dl)
@@ -0,0 +1,854 @@
# Copyright (c) 2019-present, Meta, Inc.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# First author is Simon Rouard.
import random
import typing as tp
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
import math
def create_sin_embedding(
length: int, dim: int, shift: int = 0, device="cpu", max_period=10000
):
# We aim for TBC format
assert dim % 2 == 0
pos = shift + torch.arange(length, device=device).view(-1, 1, 1)
half_dim = dim // 2
adim = torch.arange(dim // 2, device=device).view(1, 1, -1)
phase = pos / (max_period ** (adim / (half_dim - 1)))
return torch.cat(
[
torch.cos(phase),
torch.sin(phase),
],
dim=-1,
)
def create_2d_sin_embedding(d_model, height, width, device="cpu", max_period=10000):
"""
:param d_model: dimension of the model
:param height: height of the positions
:param width: width of the positions
:return: d_model*height*width position matrix
"""
if d_model % 4 != 0:
raise ValueError(
"Cannot use sin/cos positional encoding with "
"odd dimension (got dim={:d})".format(d_model)
)
pe = torch.zeros(d_model, height, width)
# Each dimension use half of d_model
d_model = int(d_model / 2)
div_term = torch.exp(
torch.arange(0.0, d_model, 2) * -(math.log(max_period) / d_model)
)
pos_w = torch.arange(0.0, width).unsqueeze(1)
pos_h = torch.arange(0.0, height).unsqueeze(1)
pe[0:d_model:2, :, :] = (
torch.sin(pos_w * div_term).transpose(0, 1).unsqueeze(1).repeat(1, height, 1)
)
pe[1:d_model:2, :, :] = (
torch.cos(pos_w * div_term).transpose(0, 1).unsqueeze(1).repeat(1, height, 1)
)
pe[d_model::2, :, :] = (
torch.sin(pos_h * div_term).transpose(0, 1).unsqueeze(2).repeat(1, 1, width)
)
pe[d_model + 1:: 2, :, :] = (
torch.cos(pos_h * div_term).transpose(0, 1).unsqueeze(2).repeat(1, 1, width)
)
return pe[None, :].to(device)
def create_sin_embedding_cape(
length: int,
dim: int,
batch_size: int,
mean_normalize: bool,
augment: bool, # True during training
max_global_shift: float = 0.0, # delta max
max_local_shift: float = 0.0, # epsilon max
max_scale: float = 1.0,
device: str = "cpu",
max_period: float = 10000.0,
):
# We aim for TBC format
assert dim % 2 == 0
pos = 1.0 * torch.arange(length).view(-1, 1, 1) # (length, 1, 1)
pos = pos.repeat(1, batch_size, 1) # (length, batch_size, 1)
if mean_normalize:
pos -= torch.nanmean(pos, dim=0, keepdim=True)
if augment:
delta = np.random.uniform(
-max_global_shift, +max_global_shift, size=[1, batch_size, 1]
)
delta_local = np.random.uniform(
-max_local_shift, +max_local_shift, size=[length, batch_size, 1]
)
log_lambdas = np.random.uniform(
-np.log(max_scale), +np.log(max_scale), size=[1, batch_size, 1]
)
pos = (pos + delta + delta_local) * np.exp(log_lambdas)
pos = pos.to(device)
half_dim = dim // 2
adim = torch.arange(dim // 2, device=device).view(1, 1, -1)
phase = pos / (max_period ** (adim / (half_dim - 1)))
return torch.cat(
[
torch.cos(phase),
torch.sin(phase),
],
dim=-1,
).float()
def get_causal_mask(length):
pos = torch.arange(length)
return pos > pos[:, None]
def get_elementary_mask(
T1,
T2,
mask_type,
sparse_attn_window,
global_window,
mask_random_seed,
sparsity,
device,
):
"""
When the input of the Decoder has length T1 and the output T2
The mask matrix has shape (T2, T1)
"""
assert mask_type in ["diag", "jmask", "random", "global"]
if mask_type == "global":
mask = torch.zeros(T2, T1, dtype=torch.bool)
mask[:, :global_window] = True
line_window = int(global_window * T2 / T1)
mask[:line_window, :] = True
if mask_type == "diag":
mask = torch.zeros(T2, T1, dtype=torch.bool)
rows = torch.arange(T2)[:, None]
cols = (
(T1 / T2 * rows + torch.arange(-sparse_attn_window, sparse_attn_window + 1))
.long()
.clamp(0, T1 - 1)
)
mask.scatter_(1, cols, torch.ones(1, dtype=torch.bool).expand_as(cols))
elif mask_type == "jmask":
mask = torch.zeros(T2 + 2, T1 + 2, dtype=torch.bool)
rows = torch.arange(T2 + 2)[:, None]
t = torch.arange(0, int((2 * T1) ** 0.5 + 1))
t = (t * (t + 1) / 2).int()
t = torch.cat([-t.flip(0)[:-1], t])
cols = (T1 / T2 * rows + t).long().clamp(0, T1 + 1)
mask.scatter_(1, cols, torch.ones(1, dtype=torch.bool).expand_as(cols))
mask = mask[1:-1, 1:-1]
elif mask_type == "random":
gene = torch.Generator(device=device)
gene.manual_seed(mask_random_seed)
mask = (
torch.rand(T1 * T2, generator=gene, device=device).reshape(T2, T1)
> sparsity
)
mask = mask.to(device)
return mask
def get_mask(
T1,
T2,
mask_type,
sparse_attn_window,
global_window,
mask_random_seed,
sparsity,
device,
):
"""
Return a SparseCSRTensor mask that is a combination of elementary masks
mask_type can be a combination of multiple masks: for instance "diag_jmask_random"
"""
from xformers.sparse import SparseCSRTensor
# create a list
mask_types = mask_type.split("_")
all_masks = [
get_elementary_mask(
T1,
T2,
mask,
sparse_attn_window,
global_window,
mask_random_seed,
sparsity,
device,
)
for mask in mask_types
]
final_mask = torch.stack(all_masks).sum(axis=0) > 0
return SparseCSRTensor.from_dense(final_mask[None])
class ScaledEmbedding(nn.Module):
def __init__(
self,
num_embeddings: int,
embedding_dim: int,
scale: float = 1.0,
boost: float = 3.0,
):
super().__init__()
self.embedding = nn.Embedding(num_embeddings, embedding_dim)
self.embedding.weight.data *= scale / boost
self.boost = boost
@property
def weight(self):
return self.embedding.weight * self.boost
def forward(self, x):
return self.embedding(x) * self.boost
class LayerScale(nn.Module):
"""Layer scale from [Touvron et al 2021] (https://arxiv.org/pdf/2103.17239.pdf).
This rescales diagonaly residual outputs close to 0 initially, then learnt.
"""
def __init__(self, channels: int, init: float = 0, channel_last=False):
"""
channel_last = False corresponds to (B, C, T) tensors
channel_last = True corresponds to (T, B, C) tensors
"""
super().__init__()
self.channel_last = channel_last
self.scale = nn.Parameter(torch.zeros(channels, requires_grad=True))
self.scale.data[:] = init
def forward(self, x):
if self.channel_last:
return self.scale * x
else:
return self.scale[:, None] * x
class MyGroupNorm(nn.GroupNorm):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
def forward(self, x):
"""
x: (B, T, C)
if num_groups=1: Normalisation on all T and C together for each B
"""
x = x.transpose(1, 2)
return super().forward(x).transpose(1, 2)
class MyTransformerEncoderLayer(nn.TransformerEncoderLayer):
def __init__(
self,
d_model,
nhead,
dim_feedforward=2048,
dropout=0.1,
activation=F.relu,
group_norm=0,
norm_first=False,
norm_out=False,
layer_norm_eps=1e-5,
layer_scale=False,
init_values=1e-4,
device=None,
dtype=None,
sparse=False,
mask_type="diag",
mask_random_seed=42,
sparse_attn_window=500,
global_window=50,
auto_sparsity=False,
sparsity=0.95,
batch_first=False,
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__(
d_model=d_model,
nhead=nhead,
dim_feedforward=dim_feedforward,
dropout=dropout,
activation=activation,
layer_norm_eps=layer_norm_eps,
batch_first=batch_first,
norm_first=norm_first,
device=device,
dtype=dtype,
)
self.sparse = sparse
self.auto_sparsity = auto_sparsity
if sparse:
if not auto_sparsity:
self.mask_type = mask_type
self.sparse_attn_window = sparse_attn_window
self.global_window = global_window
self.sparsity = sparsity
if group_norm:
self.norm1 = MyGroupNorm(int(group_norm), d_model, eps=layer_norm_eps, **factory_kwargs)
self.norm2 = MyGroupNorm(int(group_norm), d_model, eps=layer_norm_eps, **factory_kwargs)
self.norm_out = None
if self.norm_first & norm_out:
self.norm_out = MyGroupNorm(num_groups=int(norm_out), num_channels=d_model)
self.gamma_1 = (
LayerScale(d_model, init_values, True) if layer_scale else nn.Identity()
)
self.gamma_2 = (
LayerScale(d_model, init_values, True) if layer_scale else nn.Identity()
)
if sparse:
self.self_attn = MultiheadAttention(
d_model, nhead, dropout=dropout, batch_first=batch_first,
auto_sparsity=sparsity if auto_sparsity else 0,
)
self.__setattr__("src_mask", torch.zeros(1, 1))
self.mask_random_seed = mask_random_seed
def forward(self, src, src_mask=None, src_key_padding_mask=None):
"""
if batch_first = False, src shape is (T, B, C)
the case where batch_first=True is not covered
"""
device = src.device
x = src
T, B, C = x.shape
if self.sparse and not self.auto_sparsity:
assert src_mask is None
src_mask = self.src_mask
if src_mask.shape[-1] != T:
src_mask = get_mask(
T,
T,
self.mask_type,
self.sparse_attn_window,
self.global_window,
self.mask_random_seed,
self.sparsity,
device,
)
self.__setattr__("src_mask", src_mask)
if self.norm_first:
x = x + self.gamma_1(
self._sa_block(self.norm1(x), src_mask, src_key_padding_mask)
)
x = x + self.gamma_2(self._ff_block(self.norm2(x)))
if self.norm_out:
x = self.norm_out(x)
else:
x = self.norm1(
x + self.gamma_1(self._sa_block(x, src_mask, src_key_padding_mask))
)
x = self.norm2(x + self.gamma_2(self._ff_block(x)))
return x
class CrossTransformerEncoderLayer(nn.Module):
def __init__(
self,
d_model: int,
nhead: int,
dim_feedforward: int = 2048,
dropout: float = 0.1,
activation=F.relu,
layer_norm_eps: float = 1e-5,
layer_scale: bool = False,
init_values: float = 1e-4,
norm_first: bool = False,
group_norm: bool = False,
norm_out: bool = False,
sparse=False,
mask_type="diag",
mask_random_seed=42,
sparse_attn_window=500,
global_window=50,
sparsity=0.95,
auto_sparsity=None,
device=None,
dtype=None,
batch_first=False,
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.sparse = sparse
self.auto_sparsity = auto_sparsity
if sparse:
if not auto_sparsity:
self.mask_type = mask_type
self.sparse_attn_window = sparse_attn_window
self.global_window = global_window
self.sparsity = sparsity
self.cross_attn: nn.Module
self.cross_attn = nn.MultiheadAttention(
d_model, nhead, dropout=dropout, batch_first=batch_first)
# Implementation of Feedforward model
self.linear1 = nn.Linear(d_model, dim_feedforward, **factory_kwargs)
self.dropout = nn.Dropout(dropout)
self.linear2 = nn.Linear(dim_feedforward, d_model, **factory_kwargs)
self.norm_first = norm_first
self.norm1: nn.Module
self.norm2: nn.Module
self.norm3: nn.Module
if group_norm:
self.norm1 = MyGroupNorm(int(group_norm), d_model, eps=layer_norm_eps, **factory_kwargs)
self.norm2 = MyGroupNorm(int(group_norm), d_model, eps=layer_norm_eps, **factory_kwargs)
self.norm3 = MyGroupNorm(int(group_norm), d_model, eps=layer_norm_eps, **factory_kwargs)
else:
self.norm1 = nn.LayerNorm(d_model, eps=layer_norm_eps, **factory_kwargs)
self.norm2 = nn.LayerNorm(d_model, eps=layer_norm_eps, **factory_kwargs)
self.norm3 = nn.LayerNorm(d_model, eps=layer_norm_eps, **factory_kwargs)
self.norm_out = None
if self.norm_first & norm_out:
self.norm_out = MyGroupNorm(num_groups=int(norm_out), num_channels=d_model)
self.gamma_1 = (
LayerScale(d_model, init_values, True) if layer_scale else nn.Identity()
)
self.gamma_2 = (
LayerScale(d_model, init_values, True) if layer_scale else nn.Identity()
)
self.dropout1 = nn.Dropout(dropout)
self.dropout2 = nn.Dropout(dropout)
# Legacy string support for activation function.
if isinstance(activation, str):
self.activation = self._get_activation_fn(activation)
else:
self.activation = activation
if sparse:
self.cross_attn = MultiheadAttention(
d_model, nhead, dropout=dropout, batch_first=batch_first,
auto_sparsity=sparsity if auto_sparsity else 0)
if not auto_sparsity:
self.__setattr__("mask", torch.zeros(1, 1))
self.mask_random_seed = mask_random_seed
def forward(self, q, k, mask=None):
"""
Args:
q: tensor of shape (T, B, C)
k: tensor of shape (S, B, C)
mask: tensor of shape (T, S)
"""
device = q.device
T, B, C = q.shape
S, B, C = k.shape
if self.sparse and not self.auto_sparsity:
assert mask is None
mask = self.mask
if mask.shape[-1] != S or mask.shape[-2] != T:
mask = get_mask(
S,
T,
self.mask_type,
self.sparse_attn_window,
self.global_window,
self.mask_random_seed,
self.sparsity,
device,
)
self.__setattr__("mask", mask)
if self.norm_first:
x = q + self.gamma_1(self._ca_block(self.norm1(q), self.norm2(k), mask))
x = x + self.gamma_2(self._ff_block(self.norm3(x)))
if self.norm_out:
x = self.norm_out(x)
else:
x = self.norm1(q + self.gamma_1(self._ca_block(q, k, mask)))
x = self.norm2(x + self.gamma_2(self._ff_block(x)))
return x
# self-attention block
def _ca_block(self, q, k, attn_mask=None):
x = self.cross_attn(q, k, k, attn_mask=attn_mask, need_weights=False)[0]
return self.dropout1(x)
# feed forward block
def _ff_block(self, x):
x = self.linear2(self.dropout(self.activation(self.linear1(x))))
return self.dropout2(x)
def _get_activation_fn(self, activation):
if activation == "relu":
return F.relu
elif activation == "gelu":
return F.gelu
raise RuntimeError("activation should be relu/gelu, not {}".format(activation))
# ----------------- MULTI-BLOCKS MODELS: -----------------------
class CrossTransformerEncoder(nn.Module):
def __init__(
self,
dim: int,
emb: str = "sin",
hidden_scale: float = 4.0,
num_heads: int = 8,
num_layers: int = 6,
cross_first: bool = False,
dropout: float = 0.0,
max_positions: int = 1000,
norm_in: bool = True,
norm_in_group: bool = False,
group_norm: int = False,
norm_first: bool = False,
norm_out: bool = False,
max_period: float = 10000.0,
weight_decay: float = 0.0,
lr: tp.Optional[float] = None,
layer_scale: bool = False,
gelu: bool = True,
sin_random_shift: int = 0,
weight_pos_embed: float = 1.0,
cape_mean_normalize: bool = True,
cape_augment: bool = True,
cape_glob_loc_scale: list = [5000.0, 1.0, 1.4],
sparse_self_attn: bool = False,
sparse_cross_attn: bool = False,
mask_type: str = "diag",
mask_random_seed: int = 42,
sparse_attn_window: int = 500,
global_window: int = 50,
auto_sparsity: bool = False,
sparsity: float = 0.95,
):
super().__init__()
"""
"""
assert dim % num_heads == 0
hidden_dim = int(dim * hidden_scale)
self.num_layers = num_layers
# classic parity = 1 means that if idx%2 == 1 there is a
# classical encoder else there is a cross encoder
self.classic_parity = 1 if cross_first else 0
self.emb = emb
self.max_period = max_period
self.weight_decay = weight_decay
self.weight_pos_embed = weight_pos_embed
self.sin_random_shift = sin_random_shift
if emb == "cape":
self.cape_mean_normalize = cape_mean_normalize
self.cape_augment = cape_augment
self.cape_glob_loc_scale = cape_glob_loc_scale
if emb == "scaled":
self.position_embeddings = ScaledEmbedding(max_positions, dim, scale=0.2)
self.lr = lr
activation: tp.Any = F.gelu if gelu else F.relu
self.norm_in: nn.Module
self.norm_in_t: nn.Module
if norm_in:
self.norm_in = nn.LayerNorm(dim)
self.norm_in_t = nn.LayerNorm(dim)
elif norm_in_group:
self.norm_in = MyGroupNorm(int(norm_in_group), dim)
self.norm_in_t = MyGroupNorm(int(norm_in_group), dim)
else:
self.norm_in = nn.Identity()
self.norm_in_t = nn.Identity()
# spectrogram layers
self.layers = nn.ModuleList()
# temporal layers
self.layers_t = nn.ModuleList()
kwargs_common = {
"d_model": dim,
"nhead": num_heads,
"dim_feedforward": hidden_dim,
"dropout": dropout,
"activation": activation,
"group_norm": group_norm,
"norm_first": norm_first,
"norm_out": norm_out,
"layer_scale": layer_scale,
"mask_type": mask_type,
"mask_random_seed": mask_random_seed,
"sparse_attn_window": sparse_attn_window,
"global_window": global_window,
"sparsity": sparsity,
"auto_sparsity": auto_sparsity,
"batch_first": True,
}
kwargs_classic_encoder = dict(kwargs_common)
kwargs_classic_encoder.update({
"sparse": sparse_self_attn,
})
kwargs_cross_encoder = dict(kwargs_common)
kwargs_cross_encoder.update({
"sparse": sparse_cross_attn,
})
for idx in range(num_layers):
if idx % 2 == self.classic_parity:
self.layers.append(MyTransformerEncoderLayer(**kwargs_classic_encoder))
self.layers_t.append(
MyTransformerEncoderLayer(**kwargs_classic_encoder)
)
else:
self.layers.append(CrossTransformerEncoderLayer(**kwargs_cross_encoder))
self.layers_t.append(
CrossTransformerEncoderLayer(**kwargs_cross_encoder)
)
def forward(self, x, xt):
# --- BLOCK 1: Preparing x ---
B, C, Fr, T1 = x.shape
pos_emb_2d = create_2d_sin_embedding(C, Fr, T1, x.device, self.max_period) # (1, C, Fr, T1)
# Reshape to [B, Seq_Len, Channels] for the transformer
# B C Fr T1 -> B (T1 Fr) C
# The intermediate layout must be [B, T1, Fr, C]
pos_emb_2d = pos_emb_2d.permute(0, 3, 2, 1).reshape(1, T1 * Fr, C)
x = x.permute(0, 3, 2, 1).reshape(B, T1 * Fr, C)
x = self.norm_in(x)
# The batch size of 1 for pos_emb_2d will be broadcast correctly
x = x + self.weight_pos_embed * pos_emb_2d
# --- BLOCK 2: Preparing xt ---
B, C, T2 = xt.shape # B and C should be the same x and xt
xt = xt.transpose(1, 2) # B C T2 -> B T2 C
pos_emb = self._get_pos_embedding(T2, B, C, x.device) # Creates [T2, B, C]
pos_emb = pos_emb.permute(1, 0, 2) # T2 B C -> B T2 C
xt = self.norm_in_t(xt)
xt = xt + self.weight_pos_embed * pos_emb
# --- BLOCK 3: The Transformer Loop ---
for idx in range(self.num_layers):
if idx % 2 == self.classic_parity:
x = self.layers[idx](x)
xt = self.layers_t[idx](xt)
else:
old_x = x
x = self.layers[idx](x, xt)
xt = self.layers_t[idx](xt, old_x)
# --- BLOCK 4: Reshaping Outputs ---
# This is the inverse of the first operation.
# B (T1 Fr) C -> B C Fr T1
# It unflattens [B, T1*Fr, C] back to [B, T1, Fr, C]
# And then permutes back to the original [B, C, Fr, T1]
x = x.reshape(B, T1, Fr, C).permute(0, 3, 2, 1)
# Invert the transpose from the beginning
xt = xt.transpose(1, 2) # B T2 C -> B C T2
return x, xt
def _get_pos_embedding(self, T, B, C, device):
if self.emb == "sin":
shift = random.randrange(self.sin_random_shift + 1)
pos_emb = create_sin_embedding(
T, C, shift=shift, device=device, max_period=self.max_period
)
elif self.emb == "cape":
if self.training:
pos_emb = create_sin_embedding_cape(
T,
C,
B,
device=device,
max_period=self.max_period,
mean_normalize=self.cape_mean_normalize,
augment=self.cape_augment,
max_global_shift=self.cape_glob_loc_scale[0],
max_local_shift=self.cape_glob_loc_scale[1],
max_scale=self.cape_glob_loc_scale[2],
)
else:
pos_emb = create_sin_embedding_cape(
T,
C,
B,
device=device,
max_period=self.max_period,
mean_normalize=self.cape_mean_normalize,
augment=False,
)
elif self.emb == "scaled":
pos = torch.arange(T, device=device)
pos_emb = self.position_embeddings(pos)[:, None]
return pos_emb
def make_optim_group(self):
group = {"params": list(self.parameters()), "weight_decay": self.weight_decay}
if self.lr is not None:
group["lr"] = self.lr
return group
# Attention Modules
class MultiheadAttention(nn.Module):
def __init__(
self,
embed_dim,
num_heads,
dropout=0.0,
bias=True,
add_bias_kv=False,
add_zero_attn=False,
kdim=None,
vdim=None,
batch_first=False,
auto_sparsity=None,
):
super().__init__()
assert auto_sparsity is not None, "sanity check"
self.num_heads = num_heads
self.q = torch.nn.Linear(embed_dim, embed_dim, bias=bias)
self.k = torch.nn.Linear(embed_dim, embed_dim, bias=bias)
self.v = torch.nn.Linear(embed_dim, embed_dim, bias=bias)
self.attn_drop = torch.nn.Dropout(dropout)
self.proj = torch.nn.Linear(embed_dim, embed_dim, bias)
self.proj_drop = torch.nn.Dropout(dropout)
self.batch_first = batch_first
self.auto_sparsity = auto_sparsity
def forward(
self,
query,
key,
value,
key_padding_mask=None,
need_weights=True,
attn_mask=None,
average_attn_weights=True,
):
if not self.batch_first: # N, B, C
query = query.permute(1, 0, 2) # B, N_q, C
key = key.permute(1, 0, 2) # B, N_k, C
value = value.permute(1, 0, 2) # B, N_k, C
B, N_q, C = query.shape
B, N_k, C = key.shape
q = (
self.q(query)
.reshape(B, N_q, self.num_heads, C // self.num_heads)
.permute(0, 2, 1, 3)
)
q = q.flatten(0, 1)
k = (
self.k(key)
.reshape(B, N_k, self.num_heads, C // self.num_heads)
.permute(0, 2, 1, 3)
)
k = k.flatten(0, 1)
v = (
self.v(value)
.reshape(B, N_k, self.num_heads, C // self.num_heads)
.permute(0, 2, 1, 3)
)
v = v.flatten(0, 1)
if self.auto_sparsity:
assert attn_mask is None
x = dynamic_sparse_attention(q, k, v, sparsity=self.auto_sparsity)
else:
x = scaled_dot_product_attention(q, k, v, attn_mask, dropout=self.attn_drop)
x = x.reshape(B, self.num_heads, N_q, C // self.num_heads)
x = x.transpose(1, 2).reshape(B, N_q, C)
x = self.proj(x)
x = self.proj_drop(x)
if not self.batch_first:
x = x.permute(1, 0, 2)
return x, None
def scaled_query_key_softmax(q, k, att_mask):
from xformers.ops import masked_matmul
q = q / (k.size(-1)) ** 0.5
att = masked_matmul(q, k.transpose(-2, -1), att_mask)
att = torch.nn.functional.softmax(att, -1)
return att
def scaled_dot_product_attention(q, k, v, att_mask, dropout):
att = scaled_query_key_softmax(q, k, att_mask=att_mask)
att = dropout(att)
y = att @ v
return y
def _compute_buckets(x, R):
qq = torch.einsum('btf,bfhi->bhti', x, R)
qq = torch.cat([qq, -qq], dim=-1)
buckets = qq.argmax(dim=-1)
return buckets.permute(0, 2, 1).byte().contiguous()
def dynamic_sparse_attention(query, key, value, sparsity, infer_sparsity=True, attn_bias=None):
# assert False, "The code for the custom sparse kernel is not ready for release yet."
from xformers.ops import find_locations, sparse_memory_efficient_attention
n_hashes = 32
proj_size = 4
query, key, value = [x.contiguous() for x in [query, key, value]]
with torch.no_grad():
R = torch.randn(1, query.shape[-1], n_hashes, proj_size // 2, device=query.device)
bucket_query = _compute_buckets(query, R)
bucket_key = _compute_buckets(key, R)
row_offsets, column_indices = find_locations(
bucket_query, bucket_key, sparsity, infer_sparsity)
return sparse_memory_efficient_attention(
query, key, value, row_offsets, column_indices, attn_bias)
+458
View File
@@ -0,0 +1,458 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# License: MIT
import math
import typing as tp
import torch
from torch import nn
from torch.nn import functional as F
import torchaudio
from .demucs_code import capture_init, center_trim, unfold
from .CrossTransformerEncoder import LayerScale
class BLSTM(nn.Module):
"""
BiLSTM with same hidden units as input dim.
If `max_steps` is not None, input will be splitting in overlapping
chunks and the LSTM applied separately on each chunk.
"""
def __init__(self, dim, layers=1, max_steps=None, skip=False):
super().__init__()
assert max_steps is None or max_steps % 4 == 0
self.max_steps = max_steps
self.lstm = nn.LSTM(bidirectional=True, num_layers=layers, hidden_size=dim, input_size=dim)
self.linear = nn.Linear(2 * dim, dim)
self.skip = skip
def forward(self, x):
B, C, T = x.shape
y = x
framed = False
if self.max_steps is not None and T > self.max_steps:
width = self.max_steps
stride = width // 2
frames = unfold(x, width, stride)
nframes = frames.shape[2]
framed = True
x = frames.permute(0, 2, 1, 3).reshape(-1, C, width)
x = x.permute(2, 0, 1)
x = self.lstm(x)[0]
x = self.linear(x)
x = x.permute(1, 2, 0)
if framed:
out = []
frames = x.reshape(B, -1, C, width)
limit = stride // 2
for k in range(nframes):
if k == 0:
out.append(frames[:, k, :, :-limit])
elif k == nframes - 1:
out.append(frames[:, k, :, limit:])
else:
out.append(frames[:, k, :, limit:-limit])
out = torch.cat(out, -1)
out = out[..., :T]
x = out
if self.skip:
x = x + y
return x
def rescale_conv(conv, reference):
"""Rescale initial weight scale. It is unclear why it helps but it certainly does.
"""
std = conv.weight.std().detach()
scale = (std / reference)**0.5
conv.weight.data /= scale
if conv.bias is not None:
conv.bias.data /= scale
def rescale_module(module, reference):
for sub in module.modules():
if isinstance(sub, (nn.Conv1d, nn.ConvTranspose1d, nn.Conv2d, nn.ConvTranspose2d)):
rescale_conv(sub, reference)
class DConv(nn.Module):
"""
New residual branches in each encoder layer.
This alternates dilated convolutions, potentially with LSTMs and attention.
Also before entering each residual branch, dimension is projected on a smaller subspace,
e.g. of dim `channels // compress`.
"""
def __init__(self, channels: int, compress: float = 4, depth: int = 2, init: float = 1e-4,
norm=True, attn=False, heads=4, ndecay=4, lstm=False, gelu=True,
kernel=3, dilate=True):
"""
Args:
channels: input/output channels for residual branch.
compress: amount of channel compression inside the branch.
depth: number of layers in the residual branch. Each layer has its own
projection, and potentially LSTM and attention.
init: initial scale for LayerNorm.
norm: use GroupNorm.
attn: use LocalAttention.
heads: number of heads for the LocalAttention.
ndecay: number of decay controls in the LocalAttention.
lstm: use LSTM.
gelu: Use GELU activation.
kernel: kernel size for the (dilated) convolutions.
dilate: if true, use dilation, increasing with the depth.
"""
super().__init__()
assert kernel % 2 == 1
self.channels = channels
self.compress = compress
self.depth = abs(depth)
dilate = depth > 0
norm_fn: tp.Callable[[int], nn.Module]
norm_fn = lambda d: nn.Identity() # noqa
if norm:
norm_fn = lambda d: nn.GroupNorm(1, d) # noqa
hidden = int(channels / compress)
act: tp.Type[nn.Module]
if gelu:
act = nn.GELU
else:
act = nn.ReLU
self.layers = nn.ModuleList([])
for d in range(self.depth):
dilation = 2 ** d if dilate else 1
padding = dilation * (kernel // 2)
mods = [
nn.Conv1d(channels, hidden, kernel, dilation=dilation, padding=padding),
norm_fn(hidden), act(),
nn.Conv1d(hidden, 2 * channels, 1),
norm_fn(2 * channels), nn.GLU(1),
LayerScale(channels, init),
]
if attn:
mods.insert(3, LocalState(hidden, heads=heads, ndecay=ndecay))
if lstm:
mods.insert(3, BLSTM(hidden, layers=2, max_steps=200, skip=True))
layer = nn.Sequential(*mods)
self.layers.append(layer)
def forward(self, x):
for layer in self.layers:
x = x + layer(x)
return x
class LocalState(nn.Module):
"""Local state allows to have attention based only on data (no positional embedding),
but while setting a constraint on the time window (e.g. decaying penalty term).
Also a failed experiments with trying to provide some frequency based attention.
"""
def __init__(self, channels: int, heads: int = 4, nfreqs: int = 0, ndecay: int = 4):
super().__init__()
assert channels % heads == 0, (channels, heads)
self.heads = heads
self.nfreqs = nfreqs
self.ndecay = ndecay
self.content = nn.Conv1d(channels, channels, 1)
self.query = nn.Conv1d(channels, channels, 1)
self.key = nn.Conv1d(channels, channels, 1)
if nfreqs:
self.query_freqs = nn.Conv1d(channels, heads * nfreqs, 1)
if ndecay:
self.query_decay = nn.Conv1d(channels, heads * ndecay, 1)
# Initialize decay close to zero (there is a sigmoid), for maximum initial window.
self.query_decay.weight.data *= 0.01
assert self.query_decay.bias is not None # stupid type checker
self.query_decay.bias.data[:] = -2
self.proj = nn.Conv1d(channels + heads * nfreqs, channels, 1)
def forward(self, x):
B, C, T = x.shape
heads = self.heads
indexes = torch.arange(T, device=x.device, dtype=x.dtype)
# left index are keys, right index are queries
delta = indexes[:, None] - indexes[None, :]
queries = self.query(x).view(B, heads, -1, T)
keys = self.key(x).view(B, heads, -1, T)
# t are keys, s are queries
dots = torch.einsum("bhct,bhcs->bhts", keys, queries)
dots /= keys.shape[2]**0.5
if self.nfreqs:
periods = torch.arange(1, self.nfreqs + 1, device=x.device, dtype=x.dtype)
freq_kernel = torch.cos(2 * math.pi * delta / periods.view(-1, 1, 1))
freq_q = self.query_freqs(x).view(B, heads, -1, T) / self.nfreqs ** 0.5
dots += torch.einsum("fts,bhfs->bhts", freq_kernel, freq_q)
if self.ndecay:
decays = torch.arange(1, self.ndecay + 1, device=x.device, dtype=x.dtype)
decay_q = self.query_decay(x).view(B, heads, -1, T)
decay_q = torch.sigmoid(decay_q) / 2
decay_kernel = - decays.view(-1, 1, 1) * delta.abs() / self.ndecay**0.5
dots += torch.einsum("fts,bhfs->bhts", decay_kernel, decay_q)
# Kill self reference.
dots.masked_fill_(torch.eye(T, device=dots.device, dtype=torch.bool), -100)
weights = torch.softmax(dots, dim=2)
content = self.content(x).view(B, heads, -1, T)
result = torch.einsum("bhts,bhct->bhcs", weights, content)
if self.nfreqs:
time_sig = torch.einsum("bhts,fts->bhfs", weights, freq_kernel)
result = torch.cat([result, time_sig], 2)
result = result.reshape(B, -1, T)
return x + self.proj(result)
class Demucs(nn.Module):
@capture_init
def __init__(self,
sources,
# Channels
audio_channels=2,
channels=64,
growth=2.,
# Main structure
depth=6,
rewrite=True,
lstm_layers=0,
# Convolutions
kernel_size=8,
stride=4,
context=1,
# Activations
gelu=True,
glu=True,
# Normalization
norm_starts=4,
norm_groups=4,
# DConv residual branch
dconv_mode=1,
dconv_depth=2,
dconv_comp=4,
dconv_attn=4,
dconv_lstm=4,
dconv_init=1e-4,
# Pre/post processing
normalize=True,
resample=True,
# Weight init
rescale=0.1,
# Metadata
samplerate=44100,
segment=4 * 10):
"""
Args:
sources (list[str]): list of source names
audio_channels (int): stereo or mono
channels (int): first convolution channels
depth (int): number of encoder/decoder layers
growth (float): multiply (resp divide) number of channels by that
for each layer of the encoder (resp decoder)
depth (int): number of layers in the encoder and in the decoder.
rewrite (bool): add 1x1 convolution to each layer.
lstm_layers (int): number of lstm layers, 0 = no lstm. Deactivated
by default, as this is now replaced by the smaller and faster small LSTMs
in the DConv branches.
kernel_size (int): kernel size for convolutions
stride (int): stride for convolutions
context (int): kernel size of the convolution in the
decoder before the transposed convolution. If > 1,
will provide some context from neighboring time steps.
gelu: use GELU activation function.
glu (bool): use glu instead of ReLU for the 1x1 rewrite conv.
norm_starts: layer at which group norm starts being used.
decoder layers are numbered in reverse order.
norm_groups: number of groups for group norm.
dconv_mode: if 1: dconv in encoder only, 2: decoder only, 3: both.
dconv_depth: depth of residual DConv branch.
dconv_comp: compression of DConv branch.
dconv_attn: adds attention layers in DConv branch starting at this layer.
dconv_lstm: adds a LSTM layer in DConv branch starting at this layer.
dconv_init: initial scale for the DConv branch LayerScale.
normalize (bool): normalizes the input audio on the fly, and scales back
the output by the same amount.
resample (bool): upsample x2 the input and downsample /2 the output.
rescale (float): rescale initial weights of convolutions
to get their standard deviation closer to `rescale`.
samplerate (int): stored as meta information for easing
future evaluations of the model.
segment (float): duration of the chunks of audio to ideally evaluate the model on.
This is used by `demucs.apply.apply_model`.
"""
super().__init__()
self.audio_channels = audio_channels
self.sources = sources
self.kernel_size = kernel_size
self.context = context
self.stride = stride
self.depth = depth
self.resample = resample
self.channels = channels
self.normalize = normalize
self.samplerate = samplerate
self.segment = segment
self.encoder = nn.ModuleList()
self.decoder = nn.ModuleList()
self.skip_scales = nn.ModuleList()
if glu:
activation = nn.GLU(dim=1)
ch_scale = 2
else:
activation = nn.ReLU()
ch_scale = 1
if gelu:
act2 = nn.GELU
else:
act2 = nn.ReLU
in_channels = audio_channels
padding = 0
for index in range(depth):
norm_fn = lambda d: nn.Identity() # noqa
if index >= norm_starts:
norm_fn = lambda d: nn.GroupNorm(norm_groups, d) # noqa
encode = []
encode += [
nn.Conv1d(in_channels, channels, kernel_size, stride),
norm_fn(channels),
act2(),
]
attn = index >= dconv_attn
lstm = index >= dconv_lstm
if dconv_mode & 1:
encode += [DConv(channels, depth=dconv_depth, init=dconv_init,
compress=dconv_comp, attn=attn, lstm=lstm)]
if rewrite:
encode += [
nn.Conv1d(channels, ch_scale * channels, 1),
norm_fn(ch_scale * channels), activation]
self.encoder.append(nn.Sequential(*encode))
decode = []
if index > 0:
out_channels = in_channels
else:
out_channels = len(self.sources) * audio_channels
if rewrite:
decode += [
nn.Conv1d(channels, ch_scale * channels, 2 * context + 1, padding=context),
norm_fn(ch_scale * channels), activation]
if dconv_mode & 2:
decode += [DConv(channels, depth=dconv_depth, init=dconv_init,
compress=dconv_comp, attn=attn, lstm=lstm)]
decode += [nn.ConvTranspose1d(channels, out_channels,
kernel_size, stride, padding=padding)]
if index > 0:
decode += [norm_fn(out_channels), act2()]
self.decoder.insert(0, nn.Sequential(*decode))
in_channels = channels
channels = int(growth * channels)
channels = in_channels
if lstm_layers:
self.lstm = BLSTM(channels, lstm_layers)
else:
self.lstm = None
if rescale:
rescale_module(self, reference=rescale)
self.upsampler = self.downsampler = None
def valid_length(self, length):
"""
Return the nearest valid length to use with the model so that
there is no time steps left over in a convolution, e.g. for all
layers, size of the input - kernel_size % stride = 0.
Note that input are automatically padded if necessary to ensure that the output
has the same length as the input.
"""
if self.resample:
length *= 2
for _ in range(self.depth):
length = math.ceil((length - self.kernel_size) / self.stride) + 1
length = max(1, length)
for idx in range(self.depth):
length = (length - 1) * self.stride + self.kernel_size
if self.resample:
length = math.ceil(length / 2)
return int(length)
def forward(self, mix):
x = mix
length = x.shape[-1]
if self.normalize:
mono = mix.mean(dim=1, keepdim=True)
mean = mono.mean(dim=-1, keepdim=True)
std = mono.std(dim=-1, keepdim=True)
x = (x - mean) / (1e-5 + std)
else:
mean = 0
std = 1
delta = self.valid_length(length) - length
x = F.pad(x, (delta // 2, delta - delta // 2))
if self.resample:
if self.upsampler is None:
# Create the resamplers as instance attributes.
# They will be automatically moved to the correct device when model.to(device) is called.
self.upsampler = torchaudio.transforms.Resample(orig_freq=self.samplerate,
new_freq=2 * self.samplerate,
lowpass_filter_width=24).to(x.device)
self.downsampler = torchaudio.transforms.Resample(orig_freq=2 * self.samplerate,
new_freq=self.samplerate,
lowpass_filter_width=24).to(x.device)
x = self.upsampler(x)
saved = []
for encode in self.encoder:
x = encode(x)
saved.append(x)
if self.lstm:
x = self.lstm(x)
for decode in self.decoder:
skip = saved.pop(-1)
skip = center_trim(skip, x)
x = decode(x + skip)
if self.resample:
x = self.downsampler(x)
x = x * std + mean
x = center_trim(x, length)
x = x.view(x.size(0), len(self.sources), self.audio_channels, x.size(-1))
return x
def load_state_dict(self, state, strict=True):
# fix a mismatch with previous generation Demucs models.
for idx in range(self.depth):
for a in ['encoder', 'decoder']:
for b in ['bias', 'weight']:
new = f'{a}.{idx}.3.{b}'
old = f'{a}.{idx}.2.{b}'
if old in state and new not in state:
state[new] = state.pop(old)
return super().load_state_dict(state, strict=strict)
+800
View File
@@ -0,0 +1,800 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# License: MIT
"""
This code contains the spectrogram and Hybrid version of Demucs.
"""
from copy import deepcopy
import math
import typing as tp
from .wiener import wiener # From openunmix.filtering
import torch
from torch import nn
from torch.nn import functional as F
from .Demucs import DConv, rescale_module
from .demucs_code import capture_init
from .stft import spectro, ispectro
def pad1d(x: torch.Tensor, paddings: tp.Tuple[int, int], mode: str = 'constant', value: float = 0.):
"""Tiny wrapper around F.pad, just to allow for reflect padding on small input.
If this is the case, we insert extra 0 padding to the right before the reflection happen."""
x0 = x
length = x.shape[-1]
padding_left, padding_right = paddings
if mode == 'reflect':
max_pad = max(padding_left, padding_right)
if length <= max_pad:
extra_pad = max_pad - length + 1
extra_pad_right = min(padding_right, extra_pad)
extra_pad_left = extra_pad - extra_pad_right
paddings = (padding_left - extra_pad_left, padding_right - extra_pad_right)
x = F.pad(x, (extra_pad_left, extra_pad_right))
out = F.pad(x, paddings, mode, value)
assert out.shape[-1] == length + padding_left + padding_right
assert (out[..., padding_left: padding_left + length] == x0).all()
return out
class ScaledEmbedding(nn.Module):
"""
Boost learning rate for embeddings (with `scale`).
Also, can make embeddings continuous with `smooth`.
"""
def __init__(self, num_embeddings: int, embedding_dim: int,
scale: float = 10., smooth=False):
super().__init__()
self.embedding = nn.Embedding(num_embeddings, embedding_dim)
if smooth:
weight = torch.cumsum(self.embedding.weight.data, dim=0)
# when summing gaussian, overscale raises as sqrt(n), so we nornalize by that.
weight = weight / torch.arange(1, num_embeddings + 1).to(weight).sqrt()[:, None]
self.embedding.weight.data[:] = weight
self.embedding.weight.data /= scale
self.scale = scale
@property
def weight(self):
return self.embedding.weight * self.scale
def forward(self, x):
out = self.embedding(x) * self.scale
return out
class HEncLayer(nn.Module):
def __init__(self, chin, chout, kernel_size=8, stride=4, norm_groups=1, empty=False,
freq=True, dconv=True, norm=True, context=0, dconv_kw={}, pad=True,
rewrite=True, force_norm_in_last=False):
"""Encoder layer. This used both by the time and the frequency branch.
Args:
chin: number of input channels.
chout: number of output channels.
norm_groups: number of groups for group norm.
empty: used to make a layer with just the first conv. this is used
before merging the time and freq. branches.
freq: this is acting on frequencies.
dconv: insert DConv residual branches.
norm: use GroupNorm.
context: context size for the 1x1 conv.
dconv_kw: list of kwargs for the DConv class.
pad: pad the input. Padding is done so that the output size is
always the input size / stride.
rewrite: add 1x1 conv at the end of the layer.
"""
super().__init__()
norm_fn = lambda d: nn.Identity() # noqa
if norm:
norm_fn = lambda d: nn.GroupNorm(norm_groups, d) # noqa
if pad:
pad = kernel_size // 4
else:
pad = 0
klass = nn.Conv1d
self.freq = freq
self.kernel_size = kernel_size
self.stride = stride
self.empty = empty
self.norm = norm
self.pad = pad
if freq:
kernel_size = [kernel_size, 1]
stride = [stride, 1]
pad = [pad, 0]
klass = nn.Conv2d
self.conv = klass(chin, chout, kernel_size, stride, pad)
if force_norm_in_last:
# PyTorch Audio uses it on last (empty) layer
self.norm1 = norm_fn(chout)
if self.empty:
return
self.norm1 = norm_fn(chout)
self.rewrite = None
if rewrite:
self.rewrite = klass(chout, 2 * chout, 1 + 2 * context, 1, context)
self.norm2 = norm_fn(2 * chout)
self.dconv = None
if dconv:
self.dconv = DConv(chout, **dconv_kw)
def forward(self, x, inject=None):
"""
`inject` is used to inject the result from the time branch into the frequency branch,
when both have the same stride.
"""
if not self.freq and x.dim() == 4:
B, C, Fr, T = x.shape
x = x.view(B, -1, T)
if not self.freq:
le = x.shape[-1]
if not le % self.stride == 0:
x = F.pad(x, (0, self.stride - (le % self.stride)))
y = self.conv(x)
if self.empty:
return y
if inject is not None:
assert inject.shape[-1] == y.shape[-1], (inject.shape, y.shape)
if inject.dim() == 3 and y.dim() == 4:
inject = inject[:, :, None]
y = y + inject
y = F.gelu(self.norm1(y))
if self.dconv:
if self.freq:
B, C, Fr, T = y.shape
y = y.permute(0, 2, 1, 3).reshape(-1, C, T)
y = self.dconv(y)
if self.freq:
y = y.view(B, Fr, C, T).permute(0, 2, 1, 3)
if self.rewrite:
z = self.norm2(self.rewrite(y))
z = F.glu(z, dim=1)
else:
z = y
return z
class MultiWrap(nn.Module):
"""
Takes one layer and replicate it N times. each replica will act
on a frequency band. All is done so that if the N replica have the same weights,
then this is exactly equivalent to applying the original module on all frequencies.
This is a bit over-engineered to avoid edge artifacts when splitting
the frequency bands, but it is possible the naive implementation would work as well...
"""
def __init__(self, layer, split_ratios):
"""
Args:
layer: module to clone, must be either HEncLayer or HDecLayer.
split_ratios: list of float indicating which ratio to keep for each band.
"""
super().__init__()
self.split_ratios = split_ratios
self.layers = nn.ModuleList()
self.conv = isinstance(layer, HEncLayer)
assert not layer.norm
assert layer.freq
assert layer.pad
if not self.conv:
assert not layer.context_freq
for k in range(len(split_ratios) + 1):
lay = deepcopy(layer)
if self.conv:
lay.conv.padding = (0, 0)
else:
lay.pad = False
for m in lay.modules():
if hasattr(m, 'reset_parameters'):
m.reset_parameters()
self.layers.append(lay)
def forward(self, x, skip=None, length=None):
B, C, Fr, T = x.shape
ratios = list(self.split_ratios) + [1]
start = 0
outs = []
for ratio, layer in zip(ratios, self.layers):
if self.conv:
pad = layer.kernel_size // 4
if ratio == 1:
limit = Fr
frames = -1
else:
limit = int(round(Fr * ratio))
le = limit - start
if start == 0:
le += pad
frames = round((le - layer.kernel_size) / layer.stride + 1)
limit = start + (frames - 1) * layer.stride + layer.kernel_size
if start == 0:
limit -= pad
assert limit - start > 0, (limit, start)
assert limit <= Fr, (limit, Fr)
y = x[:, :, start:limit, :]
if start == 0:
y = F.pad(y, (0, 0, pad, 0))
if ratio == 1:
y = F.pad(y, (0, 0, 0, pad))
outs.append(layer(y))
start = limit - layer.kernel_size + layer.stride
else:
if ratio == 1:
limit = Fr
else:
limit = int(round(Fr * ratio))
last = layer.last
layer.last = True
y = x[:, :, start:limit]
s = skip[:, :, start:limit]
out, _ = layer(y, s, None)
if outs:
outs[-1][:, :, -layer.stride:] += (
out[:, :, :layer.stride] - layer.conv_tr.bias.view(1, -1, 1, 1))
out = out[:, :, layer.stride:]
if ratio == 1:
out = out[:, :, :-layer.stride // 2, :]
if start == 0:
out = out[:, :, layer.stride // 2:, :]
outs.append(out)
layer.last = last
start = limit
out = torch.cat(outs, dim=2)
if not self.conv and not last:
out = F.gelu(out)
if self.conv:
return out
else:
return out, None
class HDecLayer(nn.Module):
def __init__(self, chin, chout, last=False, kernel_size=8, stride=4, norm_groups=1, empty=False,
freq=True, dconv=True, norm=True, context=1, dconv_kw={}, pad=True,
context_freq=True, rewrite=True, force_norm_in_last=False):
"""
Same as HEncLayer but for decoder. See `HEncLayer` for documentation.
"""
super().__init__()
norm_fn = lambda d: nn.Identity() # noqa
if norm:
norm_fn = lambda d: nn.GroupNorm(norm_groups, d) # noqa
if pad:
pad = kernel_size // 4
else:
pad = 0
self.pad = pad
self.last = last
self.freq = freq
self.chin = chin
self.empty = empty
self.stride = stride
self.kernel_size = kernel_size
self.norm = norm
self.context_freq = context_freq
klass = nn.Conv1d
klass_tr = nn.ConvTranspose1d
if freq:
kernel_size = [kernel_size, 1]
stride = [stride, 1]
klass = nn.Conv2d
klass_tr = nn.ConvTranspose2d
self.conv_tr = klass_tr(chin, chout, kernel_size, stride)
self.norm2 = norm_fn(chout)
if self.empty:
return
self.rewrite = None
if rewrite:
if context_freq:
self.rewrite = klass(chin, 2 * chin, 1 + 2 * context, 1, context)
else:
self.rewrite = klass(chin, 2 * chin, [1, 1 + 2 * context], 1,
[0, context])
self.norm1 = norm_fn(2 * chin)
self.dconv = None
if dconv:
self.dconv = DConv(chin, **dconv_kw)
def forward(self, x, skip, length):
if self.freq and x.dim() == 3:
B, C, T = x.shape
x = x.view(B, self.chin, -1, T)
if not self.empty:
x = x + skip
if self.rewrite:
y = F.glu(self.norm1(self.rewrite(x)), dim=1)
else:
y = x
if self.dconv:
if self.freq:
B, C, Fr, T = y.shape
y = y.permute(0, 2, 1, 3).reshape(-1, C, T)
y = self.dconv(y)
if self.freq:
y = y.view(B, Fr, C, T).permute(0, 2, 1, 3)
else:
y = x
assert skip is None
z = self.norm2(self.conv_tr(y))
if self.freq:
if self.pad:
z = z[..., self.pad:-self.pad, :]
else:
z = z[..., self.pad:self.pad + length]
assert z.shape[-1] == length, (z.shape[-1], length)
if not self.last:
z = F.gelu(z)
return z, y
class HDemucs(nn.Module):
"""
Spectrogram and hybrid Demucs model.
The spectrogram model has the same structure as Demucs, except the first few layers are over the
frequency axis, until there is only 1 frequency, and then it moves to time convolutions.
Frequency layers can still access information across time steps thanks to the DConv residual.
Hybrid model have a parallel time branch. At some layer, the time branch has the same stride
as the frequency branch and then the two are combined. The opposite happens in the decoder.
Models can either use naive iSTFT from masking, Wiener filtering ([Ulhih et al. 2017]),
or complex as channels (CaC) [Choi et al. 2020]. Wiener filtering is based on
Open Unmix implementation [Stoter et al. 2019].
The loss is always on the temporal domain, by backpropagating through the above
output methods and iSTFT. This allows to define hybrid models nicely. However, this breaks
a bit Wiener filtering, as doing more iteration at test time will change the spectrogram
contribution, without changing the one from the waveform, which will lead to worse performance.
I tried using the residual option in OpenUnmix Wiener implementation, but it didn't improve.
CaC on the other hand provides similar performance for hybrid, and works naturally with
hybrid models.
This model also uses frequency embeddings are used to improve efficiency on convolutions
over the freq. axis, following [Isik et al. 2020] (https://arxiv.org/pdf/2008.04470.pdf).
Unlike classic Demucs, there is no resampling here, and normalization is always applied.
"""
@capture_init
def __init__(self,
sources,
# Channels
audio_channels=2,
channels=48,
channels_time=None,
growth=2,
# STFT
nfft=4096,
wiener_iters=0,
end_iters=0,
wiener_residual=False,
cac=True,
# Main structure
depth=6,
rewrite=True,
hybrid=True,
hybrid_old=False,
# Frequency branch
multi_freqs=None,
multi_freqs_depth=2,
freq_emb=0.2,
emb_scale=10,
emb_smooth=True,
# Convolutions
kernel_size=8,
time_stride=2,
stride=4,
context=1,
context_enc=0,
# Normalization
norm_starts=4,
norm_groups=4,
force_norm_in_last=False,
# DConv residual branch
dconv_mode=1,
dconv_depth=2,
dconv_comp=4,
dconv_attn=4,
dconv_lstm=4,
dconv_init=1e-4,
# Weight init
rescale=0.1,
# Metadata
samplerate=44100,
segment=4 * 10):
"""
Args:
sources (list[str]): list of source names.
audio_channels (int): input/output audio channels.
channels (int): initial number of hidden channels.
channels_time: if not None, use a different `channels` value for the time branch.
growth: increase the number of hidden channels by this factor at each layer.
nfft: number of fft bins. Note that changing this require careful computation of
various shape parameters and will not work out of the box for hybrid models.
wiener_iters: when using Wiener filtering, number of iterations at test time.
end_iters: same but at train time. For a hybrid model, must be equal to `wiener_iters`.
wiener_residual: add residual source before wiener filtering.
cac: uses complex as channels, i.e. complex numbers are 2 channels each
in input and output. no further processing is done before ISTFT.
depth (int): number of layers in the encoder and in the decoder.
rewrite (bool): add 1x1 convolution to each layer.
hybrid (bool): make a hybrid time/frequency domain, otherwise frequency only.
hybrid_old: some models trained for MDX had a padding bug. This replicates
this bug to avoid retraining them.
multi_freqs: list of frequency ratios for splitting frequency bands with `MultiWrap`.
multi_freqs_depth: how many layers to wrap with `MultiWrap`. Only the outermost
layers will be wrapped.
freq_emb: add frequency embedding after the first frequency layer if > 0,
the actual value controls the weight of the embedding.
emb_scale: equivalent to scaling the embedding learning rate
emb_smooth: initialize the embedding with a smooth one (with respect to frequencies).
kernel_size: kernel_size for encoder and decoder layers.
stride: stride for encoder and decoder layers.
time_stride: stride for the final time layer, after the merge.
context: context for 1x1 conv in the decoder.
context_enc: context for 1x1 conv in the encoder.
norm_starts: layer at which group norm starts being used.
decoder layers are numbered in reverse order.
norm_groups: number of groups for group norm.
dconv_mode: if 1: dconv in encoder only, 2: decoder only, 3: both.
dconv_depth: depth of residual DConv branch.
dconv_comp: compression of DConv branch.
dconv_attn: adds attention layers in DConv branch starting at this layer.
dconv_lstm: adds a LSTM layer in DConv branch starting at this layer.
dconv_init: initial scale for the DConv branch LayerScale.
rescale: weight recaling trick
"""
super().__init__()
self.cac = cac
self.wiener_residual = wiener_residual
self.audio_channels = audio_channels
self.sources = sources
self.kernel_size = kernel_size
self.context = context
self.stride = stride
self.depth = depth
self.channels = channels
self.samplerate = samplerate
self.segment = segment
self.nfft = nfft
self.hop_length = nfft // 4
self.wiener_iters = wiener_iters
self.end_iters = end_iters
self.freq_emb = None
self.hybrid = hybrid
self.hybrid_old = hybrid_old
if hybrid_old:
assert hybrid, "hybrid_old must come with hybrid=True"
if hybrid:
assert wiener_iters == end_iters
self.encoder = nn.ModuleList()
self.decoder = nn.ModuleList()
if hybrid:
self.tencoder = nn.ModuleList()
self.tdecoder = nn.ModuleList()
chin = audio_channels
chin_z = chin # number of channels for the freq branch
if self.cac:
chin_z *= 2
chout = channels_time or channels
chout_z = channels
freqs = nfft // 2
for index in range(depth):
lstm = index >= dconv_lstm
attn = index >= dconv_attn
norm = index >= norm_starts
freq = freqs > 1
stri = stride
ker = kernel_size
if not freq:
assert freqs == 1
ker = time_stride * 2
stri = time_stride
pad = True
last_freq = False
if freq and freqs <= kernel_size:
ker = freqs
pad = False
last_freq = True
kw = {
'kernel_size': ker,
'stride': stri,
'freq': freq,
'pad': pad,
'norm': norm,
'rewrite': rewrite,
'norm_groups': norm_groups,
'dconv_kw': {
'lstm': lstm,
'attn': attn,
'depth': dconv_depth,
'compress': dconv_comp,
'init': dconv_init,
'gelu': True,
}
}
kwt = dict(kw)
kwt['freq'] = 0
kwt['kernel_size'] = kernel_size
kwt['stride'] = stride
kwt['pad'] = True
kwt['force_norm_in_last'] = force_norm_in_last
kw_dec = dict(kw)
multi = False
if multi_freqs and index < multi_freqs_depth:
multi = True
kw_dec['context_freq'] = False
if last_freq:
chout_z = max(chout, chout_z)
chout = chout_z
enc = HEncLayer(chin_z, chout_z,
dconv=dconv_mode & 1, context=context_enc, **kw)
if hybrid and freq:
tenc = HEncLayer(chin, chout, dconv=dconv_mode & 1, context=context_enc,
empty=last_freq, **kwt)
self.tencoder.append(tenc)
if multi:
enc = MultiWrap(enc, multi_freqs)
self.encoder.append(enc)
if index == 0:
chin = self.audio_channels * len(self.sources)
chin_z = chin
if self.cac:
chin_z *= 2
dec = HDecLayer(chout_z, chin_z, dconv=dconv_mode & 2,
last=index == 0, context=context, **kw_dec)
if multi:
dec = MultiWrap(dec, multi_freqs)
if hybrid and freq:
tdec = HDecLayer(chout, chin, dconv=dconv_mode & 2, empty=last_freq,
last=index == 0, context=context, **kwt)
self.tdecoder.insert(0, tdec)
self.decoder.insert(0, dec)
chin = chout
chin_z = chout_z
chout = int(growth * chout)
chout_z = int(growth * chout_z)
if freq:
if freqs <= kernel_size:
freqs = 1
else:
freqs //= stride
if index == 0 and freq_emb:
self.freq_emb = ScaledEmbedding(
freqs, chin_z, smooth=emb_smooth, scale=emb_scale)
self.freq_emb_scale = freq_emb
if rescale:
rescale_module(self, reference=rescale)
def _spec(self, x):
hl = self.hop_length
nfft = self.nfft
x0 = x # noqa
if self.hybrid:
# We re-pad the signal in order to keep the property
# that the size of the output is exactly the size of the input
# divided by the stride (here hop_length), when divisible.
# This is achieved by padding by 1/4th of the kernel size (here nfft).
# which is not supported by torch.stft.
# Having all convolution operations follow this convention allow to easily
# align the time and frequency branches later on.
assert hl == nfft // 4
le = int(math.ceil(x.shape[-1] / hl))
pad = hl // 2 * 3
if not self.hybrid_old:
x = pad1d(x, (pad, pad + le * hl - x.shape[-1]), mode='reflect')
else:
x = pad1d(x, (pad, pad + le * hl - x.shape[-1]))
z = spectro(x, nfft, hl)[..., :-1, :]
if self.hybrid:
assert z.shape[-1] == le + 4, (z.shape, x.shape, le)
z = z[..., 2:2+le]
return z
def _ispec(self, z, length=None, scale=0):
hl = self.hop_length // (4 ** scale)
z = F.pad(z, (0, 0, 0, 1))
if self.hybrid:
z = F.pad(z, (2, 2))
pad = hl // 2 * 3
if not self.hybrid_old:
le = hl * int(math.ceil(length / hl)) + 2 * pad
else:
le = hl * int(math.ceil(length / hl))
x = ispectro(z, hl, length=le)
if not self.hybrid_old:
x = x[..., pad:pad + length]
else:
x = x[..., :length]
else:
x = ispectro(z, hl, length)
return x
def _magnitude(self, z):
# return the magnitude of the spectrogram, except when cac is True,
# in which case we just move the complex dimension to the channel one.
if self.cac:
B, C, Fr, T = z.shape
m = torch.view_as_real(z).permute(0, 1, 4, 2, 3)
m = m.reshape(B, C * 2, Fr, T)
else:
m = z.abs()
return m
def _mask(self, z, m):
# Apply masking given the mixture spectrogram `z` and the estimated mask `m`.
# If `cac` is True, `m` is actually a full spectrogram and `z` is ignored.
niters = self.wiener_iters
if self.cac:
B, S, C, Fr, T = m.shape
out = m.view(B, S, -1, 2, Fr, T).permute(0, 1, 2, 4, 5, 3)
out = torch.view_as_complex(out.contiguous())
return out
if self.training:
niters = self.end_iters
if niters < 0:
z = z[:, None]
return z / (1e-8 + z.abs()) * m
else:
return self._wiener(m, z, niters)
def _wiener(self, mag_out, mix_stft, niters):
# apply wiener filtering from OpenUnmix.
init = mix_stft.dtype
wiener_win_len = 300
residual = self.wiener_residual
B, S, C, Fq, T = mag_out.shape
mag_out = mag_out.permute(0, 4, 3, 2, 1)
mix_stft = torch.view_as_real(mix_stft.permute(0, 3, 2, 1))
outs = []
for sample in range(B):
pos = 0
out = []
for pos in range(0, T, wiener_win_len):
frame = slice(pos, pos + wiener_win_len)
z_out = wiener(
mag_out[sample, frame], mix_stft[sample, frame], niters,
residual=residual)
out.append(z_out.transpose(-1, -2))
outs.append(torch.cat(out, dim=0))
out = torch.view_as_complex(torch.stack(outs, 0))
out = out.permute(0, 4, 3, 2, 1).contiguous()
if residual:
out = out[:, :-1]
assert list(out.shape) == [B, S, C, Fq, T]
return out.to(init)
def forward(self, mix):
x = mix
length = x.shape[-1]
z = self._spec(mix)
mag = self._magnitude(z).to(mix.device)
x = mag
B, C, Fq, T = x.shape
# unlike previous Demucs, we always normalize because it is easier.
mean = x.mean(dim=(1, 2, 3), keepdim=True)
std = x.std(dim=(1, 2, 3), keepdim=True)
x = (x - mean) / (1e-5 + std)
# x will be the freq. branch input.
if self.hybrid:
# Prepare the time branch input.
xt = mix
meant = xt.mean(dim=(1, 2), keepdim=True)
stdt = xt.std(dim=(1, 2), keepdim=True)
xt = (xt - meant) / (1e-5 + stdt)
# okay, this is a giant mess I know...
saved = [] # skip connections, freq.
saved_t = [] # skip connections, time.
lengths = [] # saved lengths to properly remove padding, freq branch.
lengths_t = [] # saved lengths for time branch.
for idx, encode in enumerate(self.encoder):
lengths.append(x.shape[-1])
inject = None
if self.hybrid and idx < len(self.tencoder):
# we have not yet merged branches.
lengths_t.append(xt.shape[-1])
tenc = self.tencoder[idx]
xt = tenc(xt)
if not tenc.empty:
# save for skip connection
saved_t.append(xt)
else:
# tenc contains just the first conv., so that now time and freq.
# branches have the same shape and can be merged.
inject = xt
x = encode(x, inject)
if idx == 0 and self.freq_emb is not None:
# add frequency embedding to allow for non equivariant convolutions
# over the frequency axis.
frs = torch.arange(x.shape[-2], device=x.device)
emb = self.freq_emb(frs).t()[None, :, :, None].expand_as(x)
x = x + self.freq_emb_scale * emb
saved.append(x)
x = torch.zeros_like(x)
if self.hybrid:
xt = torch.zeros_like(x)
# initialize everything to zero (signal will go through u-net skips).
for idx, decode in enumerate(self.decoder):
skip = saved.pop(-1)
x, pre = decode(x, skip, lengths.pop(-1))
# `pre` contains the output just before final transposed convolution,
# which is used when the freq. and time branch separate.
if self.hybrid:
offset = self.depth - len(self.tdecoder)
if self.hybrid and idx >= offset:
tdec = self.tdecoder[idx - offset]
length_t = lengths_t.pop(-1)
if tdec.empty:
assert pre.shape[2] == 1, pre.shape
pre = pre[:, :, 0]
xt, _ = tdec(pre, None, length_t)
else:
skip = saved_t.pop(-1)
xt, _ = tdec(xt, skip, length_t)
# Let's make sure we used all stored skip connections.
assert len(saved) == 0
assert len(lengths_t) == 0
assert len(saved_t) == 0
S = len(self.sources)
x = x.view(B, S, -1, Fq, T)
x = x * std[:, None] + mean[:, None]
# to cpu as mps doesn't support complex numbers
# demucs issue #435 ##432
# NOTE: in this case z already is on cpu
# TODO: remove this when mps supports complex numbers
x_is_mps = x.device.type == "mps"
if x_is_mps:
x = x.cpu()
zout = self._mask(z, x)
x = self._ispec(zout, length)
# back to mps device
if x_is_mps:
x = x.to('mps')
if self.hybrid:
xt = xt.view(B, S, -1, length)
xt = xt * stdt[:, None] + meant[:, None]
x = xt + x
return x
+661
View File
@@ -0,0 +1,661 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# License: MIT
# First author is Simon Rouard.
"""
This code contains the spectrogram and Hybrid version of Demucs.
"""
import math
from .wiener import wiener # From openunmix.filtering
import torch
from torch import nn
from torch.nn import functional as F
from fractions import Fraction
from .Demucs import rescale_module
from .HDemucs import pad1d, ScaledEmbedding, HEncLayer, MultiWrap, HDecLayer
from .CrossTransformerEncoder import CrossTransformerEncoder
from .demucs_code import capture_init
from .stft import spectro, ispectro
class HTDemucs(nn.Module):
"""
Spectrogram and hybrid Demucs model.
The spectrogram model has the same structure as Demucs, except the first few layers are over the
frequency axis, until there is only 1 frequency, and then it moves to time convolutions.
Frequency layers can still access information across time steps thanks to the DConv residual.
Hybrid model have a parallel time branch. At some layer, the time branch has the same stride
as the frequency branch and then the two are combined. The opposite happens in the decoder.
Models can either use naive iSTFT from masking, Wiener filtering ([Ulhih et al. 2017]),
or complex as channels (CaC) [Choi et al. 2020]. Wiener filtering is based on
Open Unmix implementation [Stoter et al. 2019].
The loss is always on the temporal domain, by backpropagating through the above
output methods and iSTFT. This allows to define hybrid models nicely. However, this breaks
a bit Wiener filtering, as doing more iteration at test time will change the spectrogram
contribution, without changing the one from the waveform, which will lead to worse performance.
I tried using the residual option in OpenUnmix Wiener implementation, but it didn't improve.
CaC on the other hand provides similar performance for hybrid, and works naturally with
hybrid models.
This model also uses frequency embeddings are used to improve efficiency on convolutions
over the freq. axis, following [Isik et al. 2020] (https://arxiv.org/pdf/2008.04470.pdf).
Unlike classic Demucs, there is no resampling here, and normalization is always applied.
"""
@capture_init
def __init__(
self,
sources,
# Channels
audio_channels=2,
channels=48,
channels_time=None,
growth=2,
# STFT
nfft=4096,
wiener_iters=0,
end_iters=0,
wiener_residual=False,
cac=True,
# Main structure
depth=4,
rewrite=True,
# Frequency branch
multi_freqs=None,
multi_freqs_depth=3,
freq_emb=0.2,
emb_scale=10,
emb_smooth=True,
# Convolutions
kernel_size=8,
time_stride=2,
stride=4,
context=1,
context_enc=0,
# Normalization
norm_starts=4,
norm_groups=4,
# DConv residual branch
dconv_mode=1,
dconv_depth=2,
dconv_comp=8,
dconv_init=1e-3,
# Before the Transformer
bottom_channels=0,
# Transformer
t_layers=5,
t_emb="sin",
t_hidden_scale=4.0,
t_heads=8,
t_dropout=0.0,
t_max_positions=10000,
t_norm_in=True,
t_norm_in_group=False,
t_group_norm=False,
t_norm_first=True,
t_norm_out=True,
t_max_period=10000.0,
t_weight_decay=0.0,
t_lr=None,
t_layer_scale=True,
t_gelu=True,
t_weight_pos_embed=1.0,
t_sin_random_shift=0,
t_cape_mean_normalize=True,
t_cape_augment=True,
t_cape_glob_loc_scale=[5000.0, 1.0, 1.4],
t_sparse_self_attn=False,
t_sparse_cross_attn=False,
t_mask_type="diag",
t_mask_random_seed=42,
t_sparse_attn_window=500,
t_global_window=100,
t_sparsity=0.95,
t_auto_sparsity=False,
# ------ Particular parameters
t_cross_first=False,
# Weight init
rescale=0.1,
# Metadata
samplerate=44100,
segment=10,
use_train_segment=True,
):
"""
Args:
sources (list[str]): list of source names.
audio_channels (int): input/output audio channels.
channels (int): initial number of hidden channels.
channels_time: if not None, use a different `channels` value for the time branch.
growth: increase the number of hidden channels by this factor at each layer.
nfft: number of fft bins. Note that changing this require careful computation of
various shape parameters and will not work out of the box for hybrid models.
wiener_iters: when using Wiener filtering, number of iterations at test time.
end_iters: same but at train time. For a hybrid model, must be equal to `wiener_iters`.
wiener_residual: add residual source before wiener filtering.
cac: uses complex as channels, i.e. complex numbers are 2 channels each
in input and output. no further processing is done before ISTFT.
depth (int): number of layers in the encoder and in the decoder.
rewrite (bool): add 1x1 convolution to each layer.
multi_freqs: list of frequency ratios for splitting frequency bands with `MultiWrap`.
multi_freqs_depth: how many layers to wrap with `MultiWrap`. Only the outermost
layers will be wrapped.
freq_emb: add frequency embedding after the first frequency layer if > 0,
the actual value controls the weight of the embedding.
emb_scale: equivalent to scaling the embedding learning rate
emb_smooth: initialize the embedding with a smooth one (with respect to frequencies).
kernel_size: kernel_size for encoder and decoder layers.
stride: stride for encoder and decoder layers.
time_stride: stride for the final time layer, after the merge.
context: context for 1x1 conv in the decoder.
context_enc: context for 1x1 conv in the encoder.
norm_starts: layer at which group norm starts being used.
decoder layers are numbered in reverse order.
norm_groups: number of groups for group norm.
dconv_mode: if 1: dconv in encoder only, 2: decoder only, 3: both.
dconv_depth: depth of residual DConv branch.
dconv_comp: compression of DConv branch.
dconv_attn: adds attention layers in DConv branch starting at this layer.
dconv_lstm: adds a LSTM layer in DConv branch starting at this layer.
dconv_init: initial scale for the DConv branch LayerScale.
bottom_channels: if >0 it adds a linear layer (1x1 Conv) before and after the
transformer in order to change the number of channels
t_layers: number of layers in each branch (waveform and spec) of the transformer
t_emb: "sin", "cape" or "scaled"
t_hidden_scale: the hidden scale of the Feedforward parts of the transformer
for instance if C = 384 (the number of channels in the transformer) and
t_hidden_scale = 4.0 then the intermediate layer of the FFN has dimension
384 * 4 = 1536
t_heads: number of heads for the transformer
t_dropout: dropout in the transformer
t_max_positions: max_positions for the "scaled" positional embedding, only
useful if t_emb="scaled"
t_norm_in: (bool) norm before addinf positional embedding and getting into the
transformer layers
t_norm_in_group: (bool) if True while t_norm_in=True, the norm is on all the
timesteps (GroupNorm with group=1)
t_group_norm: (bool) if True, the norms of the Encoder Layers are on all the
timesteps (GroupNorm with group=1)
t_norm_first: (bool) if True the norm is before the attention and before the FFN
t_norm_out: (bool) if True, there is a GroupNorm (group=1) at the end of each layer
t_max_period: (float) denominator in the sinusoidal embedding expression
t_weight_decay: (float) weight decay for the transformer
t_lr: (float) specific learning rate for the transformer
t_layer_scale: (bool) Layer Scale for the transformer
t_gelu: (bool) activations of the transformer are GeLU if True, ReLU else
t_weight_pos_embed: (float) weighting of the positional embedding
t_cape_mean_normalize: (bool) if t_emb="cape", normalisation of positional embeddings
see: https://arxiv.org/abs/2106.03143
t_cape_augment: (bool) if t_emb="cape", must be True during training and False
during the inference, see: https://arxiv.org/abs/2106.03143
t_cape_glob_loc_scale: (list of 3 floats) if t_emb="cape", CAPE parameters
see: https://arxiv.org/abs/2106.03143
t_sparse_self_attn: (bool) if True, the self attentions are sparse
t_sparse_cross_attn: (bool) if True, the cross-attentions are sparse (don't use it
unless you designed really specific masks)
t_mask_type: (str) can be "diag", "jmask", "random", "global" or any combination
with '_' between: i.e. "diag_jmask_random" (note that this is permutation
invariant i.e. "diag_jmask_random" is equivalent to "jmask_random_diag")
t_mask_random_seed: (int) if "random" is in t_mask_type, controls the seed
that generated the random part of the mask
t_sparse_attn_window: (int) if "diag" is in t_mask_type, for a query (i), and
a key (j), the mask is True id |i-j|<=t_sparse_attn_window
t_global_window: (int) if "global" is in t_mask_type, mask[:t_global_window, :]
and mask[:, :t_global_window] will be True
t_sparsity: (float) if "random" is in t_mask_type, t_sparsity is the sparsity
level of the random part of the mask.
t_cross_first: (bool) if True cross attention is the first layer of the
transformer (False seems to be better)
rescale: weight rescaling trick
use_train_segment: (bool) if True, the actual size that is used during the
training is used during inference.
"""
super().__init__()
self.cac = cac
self.wiener_residual = wiener_residual
self.audio_channels = audio_channels
self.sources = sources
self.kernel_size = kernel_size
self.context = context
self.stride = stride
self.depth = depth
self.bottom_channels = bottom_channels
self.channels = channels
self.samplerate = samplerate
self.segment = segment
self.use_train_segment = use_train_segment
self.nfft = nfft
self.hop_length = nfft // 4
self.wiener_iters = wiener_iters
self.end_iters = end_iters
self.freq_emb = None
assert wiener_iters == end_iters
self.encoder = nn.ModuleList()
self.decoder = nn.ModuleList()
self.tencoder = nn.ModuleList()
self.tdecoder = nn.ModuleList()
chin = audio_channels
chin_z = chin # number of channels for the freq branch
if self.cac:
chin_z *= 2
chout = channels_time or channels
chout_z = channels
freqs = nfft // 2
for index in range(depth):
norm = index >= norm_starts
freq = freqs > 1
stri = stride
ker = kernel_size
if not freq:
assert freqs == 1
ker = time_stride * 2
stri = time_stride
pad = True
last_freq = False
if freq and freqs <= kernel_size:
ker = freqs
pad = False
last_freq = True
kw = {
"kernel_size": ker,
"stride": stri,
"freq": freq,
"pad": pad,
"norm": norm,
"rewrite": rewrite,
"norm_groups": norm_groups,
"dconv_kw": {
"depth": dconv_depth,
"compress": dconv_comp,
"init": dconv_init,
"gelu": True,
},
}
kwt = dict(kw)
kwt["freq"] = 0
kwt["kernel_size"] = kernel_size
kwt["stride"] = stride
kwt["pad"] = True
kw_dec = dict(kw)
multi = False
if multi_freqs and index < multi_freqs_depth:
multi = True
kw_dec["context_freq"] = False
if last_freq:
chout_z = max(chout, chout_z)
chout = chout_z
enc = HEncLayer(
chin_z, chout_z, dconv=dconv_mode & 1, context=context_enc, **kw
)
if freq:
tenc = HEncLayer(
chin,
chout,
dconv=dconv_mode & 1,
context=context_enc,
empty=last_freq,
**kwt
)
self.tencoder.append(tenc)
if multi:
enc = MultiWrap(enc, multi_freqs)
self.encoder.append(enc)
if index == 0:
chin = self.audio_channels * len(self.sources)
chin_z = chin
if self.cac:
chin_z *= 2
dec = HDecLayer(
chout_z,
chin_z,
dconv=dconv_mode & 2,
last=index == 0,
context=context,
**kw_dec
)
if multi:
dec = MultiWrap(dec, multi_freqs)
if freq:
tdec = HDecLayer(
chout,
chin,
dconv=dconv_mode & 2,
empty=last_freq,
last=index == 0,
context=context,
**kwt
)
self.tdecoder.insert(0, tdec)
self.decoder.insert(0, dec)
chin = chout
chin_z = chout_z
chout = int(growth * chout)
chout_z = int(growth * chout_z)
if freq:
if freqs <= kernel_size:
freqs = 1
else:
freqs //= stride
if index == 0 and freq_emb:
self.freq_emb = ScaledEmbedding(
freqs, chin_z, smooth=emb_smooth, scale=emb_scale
)
self.freq_emb_scale = freq_emb
if rescale:
rescale_module(self, reference=rescale)
transformer_channels = channels * growth ** (depth - 1)
if bottom_channels:
self.channel_upsampler = nn.Conv1d(transformer_channels, bottom_channels, 1)
self.channel_downsampler = nn.Conv1d(
bottom_channels, transformer_channels, 1
)
self.channel_upsampler_t = nn.Conv1d(
transformer_channels, bottom_channels, 1
)
self.channel_downsampler_t = nn.Conv1d(
bottom_channels, transformer_channels, 1
)
transformer_channels = bottom_channels
if t_layers > 0:
self.crosstransformer = CrossTransformerEncoder(
dim=transformer_channels,
emb=t_emb,
hidden_scale=t_hidden_scale,
num_heads=t_heads,
num_layers=t_layers,
cross_first=t_cross_first,
dropout=t_dropout,
max_positions=t_max_positions,
norm_in=t_norm_in,
norm_in_group=t_norm_in_group,
group_norm=t_group_norm,
norm_first=t_norm_first,
norm_out=t_norm_out,
max_period=t_max_period,
weight_decay=t_weight_decay,
lr=t_lr,
layer_scale=t_layer_scale,
gelu=t_gelu,
sin_random_shift=t_sin_random_shift,
weight_pos_embed=t_weight_pos_embed,
cape_mean_normalize=t_cape_mean_normalize,
cape_augment=t_cape_augment,
cape_glob_loc_scale=t_cape_glob_loc_scale,
sparse_self_attn=t_sparse_self_attn,
sparse_cross_attn=t_sparse_cross_attn,
mask_type=t_mask_type,
mask_random_seed=t_mask_random_seed,
sparse_attn_window=t_sparse_attn_window,
global_window=t_global_window,
sparsity=t_sparsity,
auto_sparsity=t_auto_sparsity,
)
else:
self.crosstransformer = None
def _spec(self, x):
hl = self.hop_length
nfft = self.nfft
x0 = x # noqa
# We re-pad the signal in order to keep the property
# that the size of the output is exactly the size of the input
# divided by the stride (here hop_length), when divisible.
# This is achieved by padding by 1/4th of the kernel size (here nfft).
# which is not supported by torch.stft.
# Having all convolution operations follow this convention allow to easily
# align the time and frequency branches later on.
assert hl == nfft // 4
le = int(math.ceil(x.shape[-1] / hl))
pad = hl // 2 * 3
x = pad1d(x, (pad, pad + le * hl - x.shape[-1]), mode="reflect")
z = spectro(x, nfft, hl)[..., :-1, :]
assert z.shape[-1] == le + 4, (z.shape, x.shape, le)
z = z[..., 2: 2 + le]
return z
def _ispec(self, z, length=None, scale=0):
hl = self.hop_length // (4**scale)
z = F.pad(z, (0, 0, 0, 1))
z = F.pad(z, (2, 2))
pad = hl // 2 * 3
le = hl * int(math.ceil(length / hl)) + 2 * pad
x = ispectro(z, hl, length=le)
x = x[..., pad: pad + length]
return x
def _magnitude(self, z):
# return the magnitude of the spectrogram, except when cac is True,
# in which case we just move the complex dimension to the channel one.
if self.cac:
B, C, Fr, T = z.shape
m = torch.view_as_real(z).permute(0, 1, 4, 2, 3)
m = m.reshape(B, C * 2, Fr, T)
else:
m = z.abs()
return m
def _mask(self, z, m):
# Apply masking given the mixture spectrogram `z` and the estimated mask `m`.
# If `cac` is True, `m` is actually a full spectrogram and `z` is ignored.
niters = self.wiener_iters
if self.cac:
B, S, C, Fr, T = m.shape
out = m.view(B, S, -1, 2, Fr, T).permute(0, 1, 2, 4, 5, 3)
out = torch.view_as_complex(out.contiguous())
return out
if self.training:
niters = self.end_iters
if niters < 0:
z = z[:, None]
return z / (1e-8 + z.abs()) * m
else:
return self._wiener(m, z, niters)
def _wiener(self, mag_out, mix_stft, niters):
# apply wiener filtering from OpenUnmix.
init = mix_stft.dtype
wiener_win_len = 300
residual = self.wiener_residual
B, S, C, Fq, T = mag_out.shape
mag_out = mag_out.permute(0, 4, 3, 2, 1)
mix_stft = torch.view_as_real(mix_stft.permute(0, 3, 2, 1))
outs = []
for sample in range(B):
pos = 0
out = []
for pos in range(0, T, wiener_win_len):
frame = slice(pos, pos + wiener_win_len)
z_out = wiener(
mag_out[sample, frame],
mix_stft[sample, frame],
niters,
residual=residual,
)
out.append(z_out.transpose(-1, -2))
outs.append(torch.cat(out, dim=0))
out = torch.view_as_complex(torch.stack(outs, 0))
out = out.permute(0, 4, 3, 2, 1).contiguous()
if residual:
out = out[:, :-1]
assert list(out.shape) == [B, S, C, Fq, T]
return out.to(init)
def valid_length(self, length: int):
"""
Return a length that is appropriate for evaluation.
In our case, always return the training length, unless
it is smaller than the given length, in which case this
raises an error.
"""
if not self.use_train_segment:
return length
training_length = int(self.segment * self.samplerate)
if training_length < length:
raise ValueError(
f"Given length {length} is longer than "
f"training length {training_length}")
return training_length
def forward(self, mix):
length = mix.shape[-1]
length_pre_pad = None
if self.use_train_segment:
if self.training:
self.segment = Fraction(mix.shape[-1], self.samplerate)
else:
training_length = int(self.segment * self.samplerate)
if mix.shape[-1] < training_length:
length_pre_pad = mix.shape[-1]
mix = F.pad(mix, (0, training_length - length_pre_pad))
z = self._spec(mix)
mag = self._magnitude(z).to(mix.device)
x = mag
B, C, Fq, T = x.shape
# unlike previous Demucs, we always normalize because it is easier.
mean = x.mean(dim=(1, 2, 3), keepdim=True)
std = x.std(dim=(1, 2, 3), keepdim=True)
x = (x - mean) / (1e-5 + std)
# x will be the freq. branch input.
# Prepare the time branch input.
xt = mix
meant = xt.mean(dim=(1, 2), keepdim=True)
stdt = xt.std(dim=(1, 2), keepdim=True)
xt = (xt - meant) / (1e-5 + stdt)
# okay, this is a giant mess I know...
saved = [] # skip connections, freq.
saved_t = [] # skip connections, time.
lengths = [] # saved lengths to properly remove padding, freq branch.
lengths_t = [] # saved lengths for time branch.
for idx, encode in enumerate(self.encoder):
lengths.append(x.shape[-1])
inject = None
if idx < len(self.tencoder):
# we have not yet merged branches.
lengths_t.append(xt.shape[-1])
tenc = self.tencoder[idx]
xt = tenc(xt)
if not tenc.empty:
# save for skip connection
saved_t.append(xt)
else:
# tenc contains just the first conv., so that now time and freq.
# branches have the same shape and can be merged.
inject = xt
x = encode(x, inject)
if idx == 0 and self.freq_emb is not None:
# add frequency embedding to allow for non equivariant convolutions
# over the frequency axis.
frs = torch.arange(x.shape[-2], device=x.device)
emb = self.freq_emb(frs).t()[None, :, :, None].expand_as(x)
x = x + self.freq_emb_scale * emb
saved.append(x)
if self.crosstransformer:
if self.bottom_channels:
b, c, f, t = x.shape
x = x.reshape(b, c, f * t) # b c f t -> b c (f t)
x = self.channel_upsampler(x)
x = x.reshape(b, x.shape[1], f, t) # b c (f t) -> b c f t
xt = self.channel_upsampler_t(xt)
x, xt = self.crosstransformer(x, xt)
if self.bottom_channels:
b, c, f, t = x.shape
x = x.reshape(b, c, f * t) # b c f t -> b c (f t)
x = self.channel_downsampler(x)
x = x.reshape(b, x.shape[1], f, t) # b c (f t) -> b c f t
xt = self.channel_downsampler_t(xt)
for idx, decode in enumerate(self.decoder):
skip = saved.pop(-1)
x, pre = decode(x, skip, lengths.pop(-1))
# `pre` contains the output just before final transposed convolution,
# which is used when the freq. and time branch separate.
offset = self.depth - len(self.tdecoder)
if idx >= offset:
tdec = self.tdecoder[idx - offset]
length_t = lengths_t.pop(-1)
if tdec.empty:
assert pre.shape[2] == 1, pre.shape
pre = pre[:, :, 0]
xt, _ = tdec(pre, None, length_t)
else:
skip = saved_t.pop(-1)
xt, _ = tdec(xt, skip, length_t)
# Let's make sure we used all stored skip connections.
assert len(saved) == 0
assert len(lengths_t) == 0
assert len(saved_t) == 0
S = len(self.sources)
x = x.view(B, S, -1, Fq, T)
x = x * std[:, None] + mean[:, None]
# to cpu as mps doesn't support complex numbers
# demucs issue #435 ##432
# NOTE: in this case z already is on cpu
# TODO: remove this when mps supports complex numbers
x_is_mps = x.device.type == "mps"
if x_is_mps:
x = x.cpu()
zout = self._mask(z, x)
if self.use_train_segment:
if self.training:
x = self._ispec(zout, length)
else:
x = self._ispec(zout, training_length)
else:
x = self._ispec(zout, length)
# back to mps device
if x_is_mps:
x = x.to("mps")
if self.use_train_segment:
if self.training:
xt = xt.view(B, S, -1, length)
else:
xt = xt.view(B, S, -1, training_length)
else:
xt = xt.view(B, S, -1, length)
xt = xt * stdt[:, None] + meant[:, None]
x = xt + x
if length_pre_pad:
x = x[..., :length_pre_pad]
return x
+379
View File
@@ -0,0 +1,379 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# Wrappers for the model and inference
import logging
import math
from seconohe.torch import model_to_target, get_offload_device
import torch
import tqdm
# ComfyUI imports
try:
import comfy.utils
with_comfy = True
except Exception:
with_comfy = False
# Local imports
from .stft import stft_chunk_process, stft_get_chunks
from ..db.load_model import load_model
from .. import NODES_NAME
from .demucs_api import apply_model, BagOfModels
from torchaudio.transforms import Fade
logger = logging.getLogger(f"{NODES_NAME}.demixer")
SAMPLE_RATE = 44100
class DemixerGeneric(object):
def __init__(self, d, device, models_dir):
super().__init__()
self.d = d
self.model_run = load_model(d, device, models_dir)
self.device = device
def set_device(self, device):
if device == self.device:
return
self.device = device
if hasattr(self.model_run, "target_device"):
self.model_run = device
# ############################################################################################################################
# MDX-Net
# ############################################################################################################################
def show_inference_parameters(d):
logger.debug("Using inference parameters:")
logger.debug(f" Frequency Bins (n_fft/2): {d['mdx_n_fft_scale_set']//2}")
logger.debug(f" Amplitude Compensation: {d['compensate']}")
class DemixerMDX(DemixerGeneric):
def __init__(self, d, device, models_dir):
super().__init__(d, device, models_dir)
show_inference_parameters(d)
self.sr = SAMPLE_RATE
self.ch = 2
def __call__(self, waveform, segments=None):
if segments is None:
# To make it compatible with Demucs
segments = 1
dim_t = (2 ** self.d['mdx_dim_t_set']) * segments
try:
# --- 1. Normalize input shape to handle both batched and non-batched data ---
if waveform.ndim == 2:
# Input is [C, samples], add a batch dimension to make it [1, C, samples]
logger.debug("Input is not batched. Adding a temporary batch dimension.")
waveform = waveform.unsqueeze(0)
input_was_batched = False
elif waveform.ndim == 3:
# Input is already batched [B, C, samples]
input_was_batched = True
else:
raise ValueError(f"Unsupported waveform shape: {waveform.shape}. Expected 2 or 3 dimensions.")
batch_size = waveform.shape[0]
logger.info("🎛️ Performing demix...")
# Lists to store the separated stems from each item in the batch
list_of_main_stems = []
list_of_complement_stems = []
# ComfyUI progress bar
progress_bar_ui = None
if with_comfy:
chunks = stft_get_chunks(waveform.shape[2], self.d['mdx_n_fft_scale_set'], segment_size=dim_t)
chunks *= batch_size
progress_bar_ui = comfy.utils.ProgressBar(chunks)
# --- 2. Iterate through the batch ---
for i, single_waveform in enumerate(waveform):
# single_waveform has shape [C, samples]
logger.debug(f"Processing item {i+1}/{batch_size}...")
# Process this single waveform
main_wav = stft_chunk_process(single_waveform, self.d, self.model_run, self.device, segment_size=dim_t,
progress_bar_ui=progress_bar_ui)
complement_wav = single_waveform - main_wav
# Add the results to our lists
list_of_main_stems.append(main_wav)
list_of_complement_stems.append(complement_wav)
# --- 3. Stack the results into single batch tensors ---
# torch.stack creates a new dimension (the batch dimension) from a list of tensors
stacked_main_stems = torch.stack(list_of_main_stems, dim=0)
stacked_complement_stems = torch.stack(list_of_complement_stems, dim=0)
# Both will now have shape [B, C, samples]
# --- 4. Denormalize output shape if original input was not batched ---
if not input_was_batched:
logger.debug("Squeezing batch dimension from output to match non-batched input.")
stacked_main_stems = stacked_main_stems.squeeze(0)
stacked_complement_stems = stacked_complement_stems.squeeze(0)
return [{'waveform': stacked_main_stems, 'sample_rate': SAMPLE_RATE, 'stem': self.d['primary_stem']},
{'waveform': stacked_complement_stems, 'sample_rate': SAMPLE_RATE, 'stem': 'Complement'}]
except Exception as e:
logger.error(f"Error during separation: {str(e)}")
raise e
# ############################################################################################################################
# Demucs
# ############################################################################################################################
def get_steps_for_demucs(model, wav, segment, shifts, overlap):
segment = segment or model.segment
segment_length = int(model.samplerate * segment)
stride = int((1 - overlap) * segment_length)
return math.ceil(wav.shape[-1] / stride) * (shifts + 1)
def separate_sources(
model: torch.nn.Module,
mix: torch.Tensor,
sample_rate: int,
segment: float = 10.0,
overlap: float = 0.1,
device: torch.device = None,
chunk_fade_shape: str = "linear",
) -> torch.Tensor:
"""
From: https://pytorch.org/audio/stable/tutorials/hybrid_demucs_tutorial.html
Apply model to a given mixture. Use fade, and add segments together in order to add model segment by segment.
Args:
segment (int): segment length in seconds
device (torch.device, str, or None): if provided, device on which to
execute the computation, otherwise `mix.device` is assumed.
When `device` is different from `mix.device`, only local computations will
be on `device`, while the entire tracks will be stored on `mix.device`.
"""
batch, channels, length = mix.shape
chunk_len = int(sample_rate * segment * (1 + overlap))
start = 0
end = chunk_len
overlap_frames = overlap * sample_rate
fade = Fade(fade_in_len=0, fade_out_len=int(overlap_frames), fade_shape=chunk_fade_shape)
chunks = math.ceil((length - overlap_frames) / chunk_len)
# Progress bars
progress_bar_console = tqdm.tqdm(total=chunks)
if with_comfy:
comfy_progress_bar = comfy.utils.ProgressBar(chunks)
final = torch.zeros(batch, len(model.sources), channels, length, device=device)
while start < length - overlap_frames:
chunk = mix[:, :, start:end]
progress_bar_console.update(1)
if with_comfy:
comfy_progress_bar.update(1)
with torch.no_grad():
out = model.forward(chunk)
out = fade(out)
final[:, :, :, start:end] += out
if start == 0:
fade.fade_in_len = int(overlap_frames)
start += int(chunk_len - overlap_frames)
else:
start += chunk_len
end += chunk_len
if end >= length:
fade.fade_out_len = 0
return final
class DemixerDemucs(DemixerGeneric):
def __init__(self, d, device, models_dir):
super().__init__(d, device, models_dir)
self.sr = self.model_run.samplerate
# Demucs code will move the model to and from the device
# The advantage is that it will be do it for the sub_models
# So here we keep it offloaded and tell Demucs code to do the work
self.model_run.target_device = get_offload_device()
def get_steps(self, wav, segment, shifts, overlap):
""" Tries to figure out how much steps we will need for inference """
if isinstance(self.model_run, BagOfModels):
total = 0
for sub_model in self.model_run.models:
steps = get_steps_for_demucs(sub_model, wav, segment, shifts, overlap)
total += steps
logger.debug(f"- Steps {steps}")
logger.debug(f"Total steps {total}")
return total
steps = get_steps_for_demucs(self.model_run, wav, segment, shifts, overlap)
logger.debug(f"Steps {steps}")
return steps
def demucs_callback(self, v):
self.comfy_progress_bar.update(1)
def __call__(self, waveform_tensor, segment=None, shifts=0, overlap=0.25):
try:
# --- 1. Normalize input shape to handle both batched and non-batched data ---
if waveform_tensor.ndim == 2:
# Input is [C, samples], add a batch dimension to make it [1, C, samples]
logger.debug("Input is not batched. Adding a temporary batch dimension.")
waveform_tensor = waveform_tensor.unsqueeze(0)
input_was_batched = False
elif waveform_tensor.ndim == 3:
# Input is already batched [B, C, samples]
input_was_batched = True
else:
raise ValueError(f"Unsupported waveform shape: {waveform_tensor.shape}. Expected 2 or 3 dimensions.")
batch_size = waveform_tensor.shape[0]
logger.info("🎛️ Performing demix...")
model = self.model_run
input_tensor_on_device = waveform_tensor.to(self.device)
# Determine the segment size
# 1. User selection
# 2. Value in the config
# 3. Auto: from each model (when None)
forced_segment = segment
if forced_segment is None:
forced_segment = self.model_run.config_segment
if forced_segment:
logger.debug(f"Using model provided segment size {forced_segment} s")
else:
logger.debug("Using default segment size")
forced_segment = None
else:
logger.debug(f"Using user provided segment size {forced_segment} s")
if self.d.get('use_demucs_pt_process', False):
# This is for the model from torchaudio, using the Demucs code we get some strange noises
# Using their example they aren't produced
model.target_device = self.device
logger.debug("Using PyTorch Audio chunking for old model")
with model_to_target(logger, model):
separated_tensors = separate_sources(
model,
input_tensor_on_device,
model.samplerate,
segment=forced_segment or 16.0,
overlap=overlap,
device=self.device,
chunk_fade_shape="half_sine"
)
else:
if with_comfy:
comfy_progress_bar = comfy.utils.ProgressBar(self.get_steps(waveform_tensor, forced_segment, shifts,
overlap))
with model_to_target(logger, model):
separated_tensors = apply_model(
model,
input_tensor_on_device,
device=self.device,
segment=forced_segment,
shifts=shifts + 1, # Shifts 0 is disabled, 1 is just one pass, 2 is 2 passes
overlap=overlap,
split=True, # Enable chunking
progress=True, # Show a progress bar in the console
callback=lambda x: (x.get('state') == 'end') and comfy_progress_bar.update(1) if with_comfy else None,
)
# Move the final result tensor back to the CPU before creating the output dicts.
# This is good practice to free up VRAM for subsequent nodes.
separated_tensors = separated_tensors.cpu()
assert batch_size == separated_tensors.shape[0]
# The output is [batch, sources, channels, samples].
# But we will separate the stems and return a batch for each stem, so we need
# [sources, batch, channels, samples].
separated_sources = separated_tensors.permute(1, 0, 2, 3)
# The model object tells us the names of the stems it produced
model_stems = model.sources.copy()
# UVR model is a 2 stems model with vocals and non_vocals, map it gracefully
if model_stems[1] == 'non_vocals':
model_stems[1] = 'other'
logger.debug(f"Model produced {len(model_stems)} stems: {model_stems}")
# Create a dictionary mapping the stem name to the audio tensor
output_map = dict(zip(model_stems, separated_sources))
# --- Silent audio for missing stems
reference_tensor = None
# Find the first valid tensor from our output to use as a shape reference
for tensor in output_map.values():
if tensor is not None and isinstance(tensor, torch.Tensor):
reference_tensor = tensor
break
# Determine the shape and sample rate for our silent audio fallback
if reference_tensor is not None:
# If we have a successful stem, use its properties
ref_batch_size, ref_channels, ref_samples = reference_tensor.shape
else:
# EDGE CASE: All stems failed. Fall back to the input audio's properties.
logger.warning("All stems failed. Using input audio shape for silence.")
# input_audio['waveform'] has shape [batch, channels, samples]
ref_batch_size, ref_channels, ref_samples = waveform_tensor.shape
# ---
# --- Gracefully create the 6 outputs ---
# Iterate through our fixed RETURN_NAMES and get the corresponding tensor.
# If a stem name doesn't exist in our output_map (e.g., 'guitar' for a 4-stem model),
# the .get() method will return None, which ComfyUI handles correctly.
final_outputs = []
for stem_name in ("vocals", "drums", "bass", "other", "guitar", "piano"):
output_tensor = output_map.get(stem_name, None)
# If the stem was produced, wrap it back into an AUDIO dict.
# If not, append None. ComfyUI handles None outputs correctly.
if output_tensor is not None:
if not input_was_batched:
output_tensor = output_tensor.squeeze(0)
output_dict = {
"waveform": output_tensor,
"sample_rate": model.samplerate,
"stem": stem_name.capitalize(),
"generated": True,
}
final_outputs.append(output_dict)
else:
logger.debug(f"Stem '{stem_name}' not produced by model. Returning silence.")
# Create a silent tensor with the correct shape, device, and dtype
silent_waveform = torch.zeros(ref_batch_size, ref_channels, ref_samples, dtype=torch.float32)
if not input_was_batched:
silent_waveform = silent_waveform.squeeze(0)
# Wrap the silent tensor in the ComfyUI AUDIO dictionary format
silent_dict = {
"waveform": silent_waveform,
"sample_rate": model.samplerate,
"stem": stem_name.capitalize(),
"generated": False,
}
final_outputs.append(silent_dict)
return tuple(final_outputs)
except Exception as e:
logger.error(f"Error during separation: {str(e)}")
raise e
def get_demixer(d, device, models_dir):
model_t = d['model_t'].lower()
if model_t == "mdx":
return DemixerMDX(d, device, models_dir)
elif model_t == "demucs":
return DemixerDemucs(d, device, models_dir)
msg = f"Unknown model type `{model_t}`"
logger.error(msg)
raise ValueError(msg)
+353
View File
@@ -0,0 +1,353 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# License: MIT
"""
Code to apply a model to a mix. It will handle chunking with overlaps and
inteprolation between chunks, as well as the `shift trick`.
"""
from concurrent.futures import ThreadPoolExecutor
import copy
import random
from threading import Lock
import typing as tp
import torch as th
from torch import nn
from torch.nn import functional as F
import tqdm
from .Demucs import Demucs
from .HDemucs import HDemucs
from .HTDemucs import HTDemucs
from .demucs_code import center_trim, DummyPoolExecutor
# AudioSeparation stuff
import logging
from .. import NODES_NAME
logger = logging.getLogger(f"{NODES_NAME}.demucs_api")
Model = tp.Union[Demucs, HDemucs, HTDemucs]
# ############################################################################################################################
# apply.py
# ############################################################################################################################
class BagOfModels(nn.Module):
def __init__(self, models: tp.List[Model],
weights: tp.Optional[tp.List[tp.List[float]]] = None,
segment: tp.Optional[float] = None):
"""
Represents a bag of models with specific weights.
You should call `apply_model` rather than calling directly the forward here for
optimal performance.
Args:
models (list[nn.Module]): list of Demucs/HDemucs models.
weights (list[list[float]]): list of weights. If None, assumed to
be all ones, otherwise it should be a list of N list (N number of models),
each containing S floats (S number of sources).
segment (None or float): overrides the `segment` attribute of each model
(this is performed inplace, be careful is you reuse the models passed).
"""
super().__init__()
assert len(models) > 0
first = models[0]
for other in models:
assert other.sources == first.sources
assert other.samplerate == first.samplerate
assert other.audio_channels == first.audio_channels
if segment is not None:
if not isinstance(other, HTDemucs) and segment > other.segment:
other.segment = segment
self.audio_channels = first.audio_channels
self.samplerate = first.samplerate
self.sources = first.sources
self.models = nn.ModuleList(models)
if weights is None:
weights = [[1. for _ in first.sources] for _ in models]
else:
assert len(weights) == len(models)
for weight in weights:
assert len(weight) == len(first.sources)
self.weights = weights
@property
def max_allowed_segment(self) -> float:
max_allowed_segment = float('inf')
for model in self.models:
if isinstance(model, HTDemucs):
max_allowed_segment = min(max_allowed_segment, float(model.segment))
return max_allowed_segment
def forward(self, x):
raise NotImplementedError("Call `apply_model` on this.")
class TensorChunk:
def __init__(self, tensor, offset=0, length=None):
total_length = tensor.shape[-1]
assert offset >= 0
assert offset < total_length
if length is None:
length = total_length - offset
else:
length = min(total_length - offset, length)
if isinstance(tensor, TensorChunk):
self.tensor = tensor.tensor
self.offset = offset + tensor.offset
else:
self.tensor = tensor
self.offset = offset
self.length = length
self.device = tensor.device
@property
def shape(self):
shape = list(self.tensor.shape)
shape[-1] = self.length
return shape
def padded(self, target_length):
delta = target_length - self.length
total_length = self.tensor.shape[-1]
assert delta >= 0
start = self.offset - delta // 2
end = start + target_length
correct_start = max(0, start)
correct_end = min(total_length, end)
pad_left = correct_start - start
pad_right = end - correct_end
out = F.pad(self.tensor[..., correct_start:correct_end], (pad_left, pad_right))
assert out.shape[-1] == target_length
return out
def tensor_chunk(tensor_or_chunk):
if isinstance(tensor_or_chunk, TensorChunk):
return tensor_or_chunk
else:
assert isinstance(tensor_or_chunk, th.Tensor)
return TensorChunk(tensor_or_chunk)
def _replace_dict(_dict: tp.Optional[dict], *subs: tp.Tuple[tp.Hashable, tp.Any]) -> dict:
if _dict is None:
_dict = {}
else:
_dict = copy.copy(_dict)
for key, value in subs:
_dict[key] = value
return _dict
def get_model_name(model, sub_model, index):
cls_name = sub_model.__class__.__name__
signatures = getattr(model, "signatures", None)
if signatures is not None:
return f"{signatures[index]} ({cls_name})"
return cls_name
def apply_model(model: tp.Union[BagOfModels, Model],
mix: tp.Union[th.Tensor, TensorChunk],
shifts: int = 1, split: bool = True,
overlap: float = 0.25, transition_power: float = 1.,
progress: bool = False, device=None,
num_workers: int = 0, segment: tp.Optional[float] = None,
pool=None, lock=None,
callback: tp.Optional[tp.Callable[[dict], None]] = None,
callback_arg: tp.Optional[dict] = None) -> th.Tensor:
"""
Apply model to a given mixture.
Args:
shifts (int): if > 0, will shift in time `mix` by a random amount between 0 and 0.5 sec
and apply the oppositve shift to the output. This is repeated `shifts` time and
all predictions are averaged. This effectively makes the model time equivariant
and improves SDR by up to 0.2 points.
split (bool): if True, the input will be broken down in 8 seconds extracts
and predictions will be performed individually on each and concatenated.
Useful for model with large memory footprint like Tasnet.
progress (bool): if True, show a progress bar (requires split=True)
device (torch.device, str, or None): if provided, device on which to
execute the computation, otherwise `mix.device` is assumed.
When `device` is different from `mix.device`, only local computations will
be on `device`, while the entire tracks will be stored on `mix.device`.
num_workers (int): if non zero, device is 'cpu', how many threads to
use in parallel.
segment (float or None): override the model segment parameter.
"""
if device is None:
device = mix.device
else:
device = th.device(device)
if pool is None:
if num_workers > 0 and device.type == 'cpu':
pool = ThreadPoolExecutor(num_workers)
else:
pool = DummyPoolExecutor()
if lock is None:
lock = Lock()
callback_arg = _replace_dict(
callback_arg, *{"model_idx_in_bag": 0, "shift_idx": 0, "segment_offset": 0}.items()
)
kwargs: tp.Dict[str, tp.Any] = {
'shifts': shifts,
'split': split,
'overlap': overlap,
'transition_power': transition_power,
'progress': progress,
'device': device,
'pool': pool,
'segment': segment,
'lock': lock,
}
out: tp.Union[float, th.Tensor]
res: tp.Union[float, th.Tensor]
if isinstance(model, BagOfModels):
# Special treatment for bag of model.
# We explicitly apply multiple times `apply_model` so that the random shifts
# are different for each model.
estimates: tp.Union[float, th.Tensor] = 0.
totals = [0.] * len(model.sources)
callback_arg["models"] = len(model.models)
for sub_model, model_weights in zip(model.models, model.weights):
kwargs["callback"] = ((
lambda d, i=callback_arg["model_idx_in_bag"]: callback(
_replace_dict(d, ("model_idx_in_bag", i))) if callback else None)
)
original_model_device = next(iter(sub_model.parameters())).device
if device != original_model_device:
m_name = get_model_name(model, sub_model, callback_arg["model_idx_in_bag"])
logger.debug(f"Moving {m_name} model from {original_model_device} to {device}")
sub_model.to(device)
res = apply_model(sub_model, mix, **kwargs, callback_arg=callback_arg)
out = res
if device != original_model_device:
logger.debug(f"Moving {m_name} model from {device} to {original_model_device}")
sub_model.to(original_model_device)
for k, inst_weight in enumerate(model_weights):
out[:, k, :, :] *= inst_weight
totals[k] += inst_weight
estimates += out
del out
callback_arg["model_idx_in_bag"] += 1
assert isinstance(estimates, th.Tensor)
for k in range(estimates.shape[1]):
estimates[:, k, :, :] /= totals[k]
return estimates
if "models" not in callback_arg:
callback_arg["models"] = 1
original_model_device = next(iter(model.parameters())).device
if device != original_model_device:
m_name = model.__class__.__name__
logger.debug(f"Moving {m_name} model from {original_model_device} to {device}")
model.to(device)
# model.eval()
assert transition_power >= 1, "transition_power < 1 leads to weird behavior."
batch, channels, length = mix.shape
if shifts:
kwargs['shifts'] = 0
max_shift = int(0.5 * model.samplerate)
mix = tensor_chunk(mix)
assert isinstance(mix, TensorChunk)
padded_mix = mix.padded(length + 2 * max_shift)
out = 0.
for shift_idx in range(shifts):
offset = random.randint(0, max_shift)
shifted = TensorChunk(padded_mix, offset, length + max_shift - offset)
kwargs["callback"] = (
(lambda d, i=shift_idx: callback(_replace_dict(d, ("shift_idx", i)))
if callback else None)
)
res = apply_model(model, shifted, **kwargs, callback_arg=callback_arg)
shifted_out = res
out += shifted_out[..., max_shift - offset:]
out /= shifts
assert isinstance(out, th.Tensor)
return out
elif split:
kwargs['split'] = False
out = th.zeros(batch, len(model.sources), channels, length, device=mix.device)
sum_weight = th.zeros(length, device=mix.device)
if segment is None: # or isinstance(model, HTDemucs):
segment = model.segment
logger.debug(f"Default model segment is: {segment} s")
assert segment is not None and segment > 0.
segment_length: int = int(model.samplerate * segment)
stride = int((1 - overlap) * segment_length)
offsets = range(0, length, stride)
# scale = float(format(stride / model.samplerate, ".2f"))
# We start from a triangle shaped weight, with maximal weight in the middle
# of the segment. Then we normalize and take to the power `transition_power`.
# Large values of transition power will lead to sharper transitions.
weight = th.cat([th.arange(1, segment_length // 2 + 1, device=device),
th.arange(segment_length - segment_length // 2, 0, -1, device=device)])
assert len(weight) == segment_length
# If the overlap < 50%, this will translate to linear transition when
# transition_power is 1.
weight = (weight / weight.max())**transition_power
futures = []
for offset in offsets:
chunk = TensorChunk(mix, offset, segment_length)
# Add more information for progress
future = pool.submit(apply_model, model, chunk, **kwargs, callback_arg=callback_arg,
callback=(lambda d, i=offset: callback(_replace_dict(d, ("segment_offset", i)))
if callback else None))
futures.append((future, offset))
offset += segment_length
if progress:
futures = tqdm.tqdm(futures, unit='seconds') # unit_scale=scale, makes a mess
for future, offset in futures:
try:
chunk_out = future.result() # type: th.Tensor
except Exception:
pool.shutdown(wait=True, cancel_futures=True)
raise
chunk_length = chunk_out.shape[-1]
out[..., offset:offset + segment_length] += (
weight[:chunk_length] * chunk_out).to(mix.device)
sum_weight[offset:offset + segment_length] += weight[:chunk_length].to(mix.device)
assert sum_weight.min() > 0
out /= sum_weight
assert isinstance(out, th.Tensor)
return out
else:
valid_length: int
if isinstance(model, HTDemucs) and segment is not None:
valid_length = int(segment * model.samplerate)
# Inform to the object that we are not using the training length
model.use_train_segment = False
elif hasattr(model, 'valid_length'):
valid_length = model.valid_length(length) # type: ignore
else:
valid_length = length
mix = tensor_chunk(mix)
assert isinstance(mix, TensorChunk)
padded_mix = mix.padded(valid_length).to(device)
with lock:
if callback is not None:
callback(_replace_dict(callback_arg, ("state", "start"))) # type: ignore
with th.no_grad():
out = model(padded_mix)
with lock:
if callback is not None:
callback(_replace_dict(callback_arg, ("state", "end"))) # type: ignore
assert isinstance(out, th.Tensor)
return center_trim(out, length)
+94
View File
@@ -0,0 +1,94 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# License: MIT
# Misc stuff used by Demucs
import functools
import math
import torch
from torch.nn import functional as F
import typing as tp
# ############################################################################################################################
# utils.py
# ############################################################################################################################
def unfold(a, kernel_size, stride):
"""Given input of size [*X, T], output Tensor of size [*X, F, K]
with K the kernel size, by extracting frames with the given stride.
This will pad the input so that `F = ceil(T / K)`.
see https://github.com/pytorch/pytorch/issues/60466
"""
*shape, length = a.shape
n_frames = math.ceil(length / stride)
tgt_length = (n_frames - 1) * stride + kernel_size
a = F.pad(a, (0, tgt_length - length))
strides = list(a.stride())
assert strides[-1] == 1, 'data should be contiguous'
strides = strides[:-1] + [stride, 1]
return a.as_strided([*shape, n_frames, kernel_size], strides)
def center_trim(tensor: torch.Tensor, reference: tp.Union[torch.Tensor, int]):
"""
Center trim `tensor` with respect to `reference`, along the last dimension.
`reference` can also be a number, representing the length to trim to.
If the size difference != 0 mod 2, the extra sample is removed on the right side.
"""
ref_size: int
if isinstance(reference, torch.Tensor):
ref_size = reference.size(-1)
else:
ref_size = reference
delta = tensor.size(-1) - ref_size
if delta < 0:
raise ValueError("tensor must be larger than reference. " f"Delta is {delta}.")
if delta:
tensor = tensor[..., delta // 2:-(delta - delta // 2)]
return tensor
class DummyPoolExecutor:
class DummyResult:
def __init__(self, func, *args, **kwargs):
self.func = func
self.args = args
self.kwargs = kwargs
def result(self):
return self.func(*self.args, **self.kwargs)
def __init__(self, workers=0):
pass
def submit(self, func, *args, **kwargs):
return DummyPoolExecutor.DummyResult(func, *args, **kwargs)
def shutdown(self, wait=True, cancel_futures=True):
return
def __enter__(self):
return self
def __exit__(self, exc_type, exc_value, exc_tb):
return
# ############################################################################################################################
# state.py
# ############################################################################################################################
def capture_init(init):
@functools.wraps(init)
def __init__(self, *args, **kwargs):
self._init_args_kwargs = (args, kwargs)
init(self, *args, **kwargs)
return __init__
+322
View File
@@ -0,0 +1,322 @@
from fractions import Fraction
import logging
from .. import NODES_NAME
rlogger = logging.getLogger(f"{NODES_NAME}.demucs_log")
class DemucsModelInfo(object):
def __init__(self, index, klass_name: str, kwargs: dict, logger, weights, extra=False, sig=None):
super().__init__()
self.index = index
self.kwargs = kwargs
self.logger = logger
self.extra = extra
self.extra_indent = ""
self.weights = weights
self.signature = sig
# Always log these fundamental parameters
sr = kwargs.get('samplerate', 44100)
if sr != 44100:
rlogger.warning("Model not configured for 44.1 kHz sample rate")
a_ch = kwargs.get('audio_channels', 2)
if a_ch != 2:
rlogger.warning("Model not configured for stereo")
if klass_name == "HTDemucs":
self.htdemucs()
elif klass_name == "HDemucs":
self.hdemucs()
elif klass_name == "Demucs":
self.demucs()
else:
logger.warning(f"No specific logger for model class: {klass_name}. "
"Displaying raw kwargs.")
for key, value in kwargs.items():
logger(f" - {key}: {value}")
def get(self, key, default=None):
return self.kwargs.get(key, default)
def log_type(self, name):
start = "" if self.index < 0 else f"{self.index+1}. "
msg = f" {start}Type: {name}"
if self.signature:
msg += f" [{self.signature}]"
self.logger(msg)
def _log_param(self, param_name, default, description="", unit="", indent=" ", can_skip=False):
"""
Logs a parameter if its value is different from the default, or if it's a key parameter.
Args:
param_name (str): The name of the parameter to check.
default: The default value for this parameter.
description (str): A user-friendly description of the parameter.
unit (str): An optional unit to display after the value (e.g., 'Hz').
indent (str): The indentation string for the log message.
"""
value = self.get(param_name, default)
# We log if the value is not the default, or if it's a fundamental parameter.
is_default = (value == default)
is_important = param_name in ['sources', 'segment']
if is_default and not is_important and not self.extra:
return None
if can_skip and is_default:
return None
if param_name == 'sources' and self.weights:
value = [s if w == 1.0 else ('' if not w else f'{w}*{s}') for s, w in zip(value, self.weights)]
if unit == '%':
value *= 100
unit_str = f" {unit}" if unit else ""
desc_str = description if description else param_name.capitalize().replace('_', ' ')
indent += self.extra_indent
value_str = f"{value.numerator}/{value.denominator}" if isinstance(value, Fraction) else str(value)
n = (39 - len(desc_str) - len(value_str) - len(unit_str))*" "
return f" {indent}- {desc_str}: {value_str}{unit_str} {n}({param_name})"
def log_param(self, param_name, default, description="", unit="", indent=" ", can_skip=False):
res = self._log_param(param_name, default, description, unit, indent, can_skip)
if res is not None:
self.logger(res)
def add(self, param_name, default, description="", unit="", indent=" ", can_skip=False):
res = self._log_param(param_name, default, description, unit, indent, can_skip)
if res is not None:
self.params.append(res)
def reset(self):
self.params = []
def sub_section(self, name):
self.logger(f" {name}:")
def section(self, name):
self.logger(" " + "-" * 40)
self.sub_section(name)
def flush(self, name, is_sub=False):
if self.params:
if is_sub:
self.sub_section(self.extra_indent + name)
else:
self.section(name)
for p in self.params:
self.logger(p)
def structure(self, ch=64, depth=6, with_lstm=False, with_ch_tm=False):
self.reset()
self.add('channels', ch, "Initial hidden channels")
self.add('depth', depth, "Number of U-Net layers")
self.add('growth', 2.0, "Channel growth factor per layer")
self.add('rewrite', True, "Use 1x1 convolutions in blocks")
if with_lstm:
self.add('lstm_layers', 0, "Number of main LSTM layers", can_skip=True)
if with_ch_tm:
self.add('channels_time', None, "Specific channels for time branch", can_skip=True)
self.flush("Structure")
def convolutions(self, advanced=False):
self.reset()
self.add('kernel_size', 8)
self.add('stride', 4)
if advanced:
self.add('time_stride', 2, "Final time layer stride")
self.add('context', 1, "Decoder context window size")
if advanced:
self.add('context_enc', 0, "Encoder context window size")
self.flush("Convolutions")
def normalization(self):
self.reset()
self.add('norm_starts', 4, "Start at layer")
self.add('norm_groups', 4, "Number of groups")
self.flush("Normalization")
def dconv(self, full=True):
if self.get('dconv_mode', 1) <= 0:
return
self.reset()
where = ['', 'In encoder', 'In decoder', 'In encoder and decoder'][self.get('dconv_mode', 1)]
self.add('dconv_mode', 1, where)
self.add('dconv_depth', 2, "Number of layers in DConv branch")
if full:
comp = 4
init = 1e-4
else:
comp = 8
init = 1e-3
self.add('dconv_comp', comp, "Channel compression factor")
self.add('dconv_init', init, "Initial scale")
if full:
self.add('dconv_attn', 4, "Layer to start attention in DConv")
self.add('dconv_lstm', 4, "Layer to start LSTM in DConv")
self.flush("DConv Residual Branch")
def stft(self):
self.reset()
self.add('nfft', 4096, "Frequency Bins")
# Decode the method
cac = self.get('cac')
niters = self.get('wiener_iters', 0)
if cac:
zout = "Complex as Channels (CaC)"
elif niters >= 0:
zout = "Wiener filtering"
else:
zout = "Naive iSTFT from masking"
self.add('___', zout, "Framework")
self.add('cac', True, "Use Complex as Channels")
if not self.get('cac', True):
self.add('wiener_iters', 0, "Wiener filter iterations")
self.flush("STFT")
def freq_branch(self):
self.reset()
def_ratio = None
if self.get('multi_freqs') == []:
def_ratio = []
self.add('multi_freqs', def_ratio, "Ratios for frequency band splitting")
if self.get('multi_freqs'):
self.add('multi_freqs_depth', 2, "Layers to apply frequency splitting")
self.add('freq_emb', 0.2, "Frequency embedding weight")
if self.get('freq_emb'):
indent = " "
self.add('emb_scale', 10, "Scale", indent=indent)
self.add('emb_smooth', True, "Smooth", indent=indent)
self.flush("Frequency Branch")
def demucs(self):
"""Logs the parameters for the original Demucs class."""
self.log_type("Classic Waveform Demucs (Demucs)")
self.log_param('sources', [], "Target source names")
self.log_param('segment', 40, "Segment size", unit="s")
# --- Structure & Channels ---
self.structure(ch=64, depth=6, with_lstm=True)
# --- Convolutions ---
self.convolutions()
self.reset()
self.add('gelu', True, "GeLU (not ReLU)")
if self.get('rewrite', True):
self.add('glu', True, "GLU in 1x1 rewrite (not ReLU)")
self.flush("Activations")
# --- Normalization ---
self.normalization()
# --- DConv Residual Branch ---
self.dconv()
# --- Pre/Post Processing ---
self.reset()
self.add('resample', True, "Use 2x resampling")
self.add('normalize', True, "Normalize audio on-the-fly")
self.flush("Processing")
def hdemucs(self):
"""Logs the parameters for the HDemucs (Hybrid Spectrogram/Waveform) class."""
self.log_type("Hybrid Demucs (Spectrogram + Waveform) (HDemucs)")
self.log_param('sources', [], "Target source names")
self.log_param('segment', 40, "Segment size", unit="s")
# --- Structure & Channels ---
self.structure(ch=48, depth=6, with_ch_tm=True)
# --- STFT & Spectrogram ---
self.stft()
# --- Frequency Branch ---
self.freq_branch()
# --- Convolutions ---
self.convolutions(advanced=True)
# --- Normalization ---
self.normalization()
# --- DConv Residual Branch (defaults are different from Demucs) ---
self.dconv()
def htdemucs(self):
"""Logs the parameters for the HTDemucs (Hybrid Transformer) class."""
self.log_type("Hybrid Transformer Demucs (HTDemucs)")
self.log_param('sources', [], "Target source names")
self.log_param('segment', 10, "Segment size", unit="s")
# --- Structure & Channels (defaults are different from HDemucs) ---
self.structure(ch=48, depth=4)
# --- STFT & Spectrogram ---
self.stft()
# --- Frequency Branch ---
self.freq_branch()
# --- Convolutions ---
self.convolutions(advanced=True)
# --- Normalization ---
self.normalization()
# --- DConv (defaults are different) ---
self.dconv(full=False)
# --- Transformer Block ---
if self.get('t_layers', 5) > 0:
self.extra_indent = " "
# --- Main Transformer ---
self.reset()
if self.get('bottom_channels', 0):
self.add('bottom_channels', 0, "Channels forced to")
self.add('t_hidden_scale', 4.0, "Hidden scale")
self.add('t_layers', 5, "Number of transformer layers")
self.add('t_heads', 8, "Number of attention heads")
self.add('t_dropout', 0.0, "Dropout")
self.flush("Transformer")
# --- Positional Embeddings ---
self.reset()
self.add('t_emb', 'sin', "Type")
self.add('t_weight_pos_embed', 1.0, "Weight", can_skip=True)
t_emb = self.get('t_emb', 'sin')
if t_emb == 'scaled':
self.add('t_max_positions', 10000, "Max positions")
elif t_emb == 'sin':
self.add('t_max_period', 10000.0, "Max period")
self.add('t_sin_random_shift', 0, "Random shift", can_skip=True)
elif t_emb == 'cape':
self.add('t_cape_mean_normalize', True, "Cape normalize")
self.add('t_cape_glob_loc_scale', [5000.0, 1.0, 1.4], "Cape params")
if self.get('t_cape_augment', True):
rlogger.warning("t_cape_augment is True in loaded model, should be False for inference.")
self.flush("Positional Embeddings", is_sub=True)
# --- Transformer Normalization ---
self.reset()
self.add('t_norm_first', True, "Before attention/FFN")
self.add('t_norm_in', True, "Before pos. embedding")
if self.get('t_norm_in', True):
self.add('t_norm_in_group', False, "On all timesteps")
self.add('t_group_norm', False, "Of encoder on all timesteps")
self.add('t_norm_out', True, "GroupNorm at end of layers")
self.flush("Normalization", is_sub=True)
# --- Transformer Misc ---
self.reset()
self.add('t_cross_first', False, "Cross-attention is the first layer")
self.add('t_layer_scale', True, "Layer scale")
self.add('t_gelu', True, "GeLU (not ReLU)")
self.flush("Various", is_sub=True)
# --- Sparsity ---
# Log sparsity details only if sparse attention is enabled
self.reset()
is_sparse = self.get('t_sparse_self_attn', False)
self.add('t_sparse_self_attn', False, "Use sparse self-attention")
if is_sparse:
self.add('t_sparse_cross_attn', False, "Sparse cross-attention")
self.add('t_auto_sparsity', False, "Automatic sparsity")
auto_sparsity = self.get('t_auto_sparsity', False)
if not auto_sparsity:
self.add('t_mask_type', 'diag', "Masking pattern")
self.add('t_mask_random_seed', 42, "Mask seed")
mask_t = self.get('t_mask_type', 'diag')
if 'diag' in mask_t:
self.add('t_sparse_attn_window', 500, "Window size")
if 'global' in mask_t:
self.add('t_global_window', 100, "Window size")
if 'random' in mask_t:
self.add('t_sparsity', 0.95, "Sparsity for random mask", unit="%")
self.flush("Sparsity", is_sub=True)
# Training only
# self.add(logger, kwargs, 't_weight_decay', 0.0, "Weight decay", extra=False)
# self.add(logger, kwargs, 't_lr', None, "Learning rate", extra=False)
# self.add(logger, kwargs, 't_cape_augment', True, "Learning rate", extra=False)
# self.add(logger, kwargs, 'rescale', 0.1, "Rescale trick", extra=False)
self.extra_indent = ""
+166
View File
@@ -0,0 +1,166 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# Helper to get a model from the correct class
import importlib
import json
import logging
from safetensors import safe_open
from seconohe.logger import get_debug_level
from .MDX_Net import MDX_Net
from .. import NODES_NAME
from ..utils.misc import json_object_hook
# Demucs class imports
from .demucs_api import BagOfModels
from .demucs_log_helper import DemucsModelInfo
logger = logging.getLogger(f"{NODES_NAME}.get_model")
def get_metadata(file_path, d=None):
""" Read the metadata from a safetensors file """
logger.debug(f"Reading metadata from {file_path}")
metadata = {}
with safe_open(file_path, framework="pt", device="cpu") as f:
metadata = f.metadata()
if not metadata:
raise ValueError(f"Could not read metadata from safetensors file: {file_path}")
if d is None:
return metadata
# Is this a child model?
parent = d.get('parent')
if parent:
# Ok, this is a child model changing details of a parent model
# Currently used by Demucs models to create simplified versions of the same model
metadata['is_bag_of_models'] = d.get('is_bag_of_models', 'false')
try:
metadata['signatures'] = d['signatures']
except KeyError:
logger.error("Child model without signatures")
raise
metadata['segment'] = d.get('segment', '0')
return metadata
def get_hyperparameter(metadata, parameter, as_type, default=None, warn_diff=True):
value = metadata.get(parameter)
if value is None:
if default is None:
raise ValueError(f"Missing `{parameter}` hyperparameter")
return default
if as_type == "int":
value = int(value)
if warn_diff and value != default:
logger.warning(f"Hyperparameter mismatch: database = {default}, metadata = {value}")
return value
def get_mdx_model(d):
""" Create an MDX_Net object with the specified parameters """
# Check the file is consistent we our data base
metadata = get_metadata(d['model_path'], d)
dim_f = get_hyperparameter(metadata, 'mdx_dim_f_set', "int", d['mdx_dim_f_set'])
channels = get_hyperparameter(metadata, 'channels', "int", d['channels'])
stages = get_hyperparameter(metadata, 'stages', "int", d['stages'])
# Create a class with this parameters
return MDX_Net(dim_f=dim_f, ch=channels, num_stages=stages)
def get_model_path(d):
parent = d.get("parent")
if parent is None:
return d.get('model_path')
return parent.get('model_path')
def get_demucs_model(d):
""" Create a Demucs, HDemucs (Hybrid) or HTDemucs (Hybrid Transformer) object.
All metadata comes from the safetensors """
file_path = get_model_path(d)
# 1. First, open the file safely to read only the metadata header.
metadata = get_metadata(file_path, d)
is_bag = json.loads(metadata.get('is_bag_of_models', 'false'))
signatures = json.loads(metadata['signatures'])
sub_models = []
for sig in set(signatures): # Use set to only instantiate each unique architecture once
model_meta_str = metadata.get(sig)
if not model_meta_str:
raise ValueError(f"Metadata for signature '{sig}' not found in safetensors file.")
# Use the object_hook here to reconstruct Fraction objects automatically
model_meta = json.loads(model_meta_str, object_hook=json_object_hook)
class_module, class_name = model_meta['class_module'], model_meta['class_name']
args, kwargs = model_meta['args'], model_meta['kwargs']
logger.debug(f" - Reconstructing architecture for '{sig}': {class_module}.{class_name}")
assert '.' not in class_name, "Security check failed, won't import a file outside my directory"
# Import the class from a module in this dir with the same name as the class
local_class_module = '.'.join(__name__.split('.')[:-1]) + "." + class_name
logger.debug(f" - Redirecting class: {class_module} -> {local_class_module}")
module = importlib.import_module(local_class_module)
klass = getattr(module, class_name)
instance = klass(*args, **kwargs)
instance._signature = sig
instance._metadata = model_meta
sub_models.append({'sig': sig, 'model': instance})
# Create a mapping from signature to model instance
model_map = {m['sig']: m['model'] for m in sub_models}
# Re-order the models to match the YAML's signature list
ordered_models = [model_map[sig] for sig in signatures]
segment = float(metadata.get('segment', '0'))
if not is_bag:
weights = None
final_model = ordered_models[0]
final_model.signatures = signatures
else:
logger.debug("Rebuilding BagOfModels container...")
weights = json.loads(metadata.get('weights', 'null'))
final_model = BagOfModels(ordered_models, weights=weights, segment=segment)
final_model.signatures = signatures
final_model.config_segment = segment
debug_level = get_debug_level(logger)
if debug_level >= 1:
# Show some information of the resulting model
logger.debug("Model information:")
logger.debug(f"- Total models {len(ordered_models)}")
if weights is not None and len(weights) != len(ordered_models):
raise ValueError(f"Invalid weights for {len(ordered_models)} models: {weights}")
for n, m in enumerate(ordered_models):
tp = m.__class__.__name__
w = weights[n] if weights else None
if n == 0:
ref_sources = m.sources
else:
if m.sources != ref_sources:
logger.error("The sub-model outputs doesn't match")
if w and len(m.sources) != len(w):
raise ValueError(f"Invalid {w} weights for {m.sources} sources")
num = -1 if len(ordered_models) == 1 else n
DemucsModelInfo(num, tp, m._metadata['kwargs'], logger.debug, w, extra=debug_level > 1, sig=m._signature)
return final_model
def get_model(d):
model_t = d['model_t'].lower()
if model_t == "mdx":
return get_mdx_model(d)
elif model_t == "demucs":
return get_demucs_model(d)
msg = f"Unknown model type `{model_t}`"
logger.error(msg)
raise ValueError(msg)
@@ -1,11 +1,11 @@
# Short-Time Fourier Transform (STFT).
import logging
import numpy as np
from seconohe.torch import model_to_target
import torch
from tqdm import tqdm
# Local imports
from ..utils.misc import NODES_NAME
from ..utils.torch import model_to_target
from .. import NODES_NAME
logger = logging.getLogger(f"{NODES_NAME}.stft")
@@ -126,9 +126,8 @@ def stft_chunk_process(waveform, d, model_run, device, segment_size=256, hop_len
total_chunks = 1 + (mixture.shape[1] - chunk_size + step - 1) // step
logger.info(f"⚙️ Processing {total_chunks} chunks...")
model_run.target_device = device
with model_to_target(model_run):
with model_to_target(logger, model_run):
for i in tqdm(range(0, mixture.shape[1] - chunk_size + 1, step)):
start = i
end = i + chunk_size
@@ -162,3 +161,51 @@ def stft_chunk_process(waveform, d, model_run, device, segment_size=256, hop_len
# Convert final result back to a torch tensor for saving
return torch.from_numpy(main_wav_np)
# ############################################################################################################################
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
# License: MIT
#
# Convenience wrapper to perform STFT and iSTFT
# ############################################################################################################################
def spectro(x, n_fft=512, hop_length=None, pad=0):
*other, length = x.shape
x = x.reshape(-1, length)
is_mps = x.device.type == 'mps'
if is_mps:
x = x.cpu()
z = torch.stft(x,
n_fft * (1 + pad),
hop_length or n_fft // 4,
window=torch.hann_window(n_fft).to(x),
win_length=n_fft,
normalized=True,
center=True,
return_complex=True,
pad_mode='reflect')
_, freqs, frame = z.shape
return z.view(*other, freqs, frame)
def ispectro(z, hop_length=None, length=None, pad=0):
*other, freqs, frames = z.shape
n_fft = 2 * freqs - 2
z = z.view(-1, freqs, frames)
win_length = n_fft // (1 + pad)
is_mps = z.device.type == 'mps'
if is_mps:
z = z.cpu()
x = torch.istft(z,
n_fft,
hop_length,
window=torch.hann_window(win_length).to(z.real),
win_length=win_length,
normalized=True,
length=length,
center=True)
_, length = x.shape
return x.view(*other, length)
+499
View File
@@ -0,0 +1,499 @@
# Authors F.R. Stoter and S. Uhlich and A. Liutkus and Y. Mitsufuji
# Project: Open-Unmix - A Reference Implementation for Music Source Separation
# License: MIT
# Site: https://github.com/sigsep/open-unmix-pytorch/tree/master/openunmix
from typing import Optional
import torch
def atan2(y, x):
r"""Element-wise arctangent function of y/x.
Returns a new tensor with signed angles in radians.
It is an alternative implementation of torch.atan2
Args:
y (Tensor): First input tensor
x (Tensor): Second input tensor [shape=y.shape]
Returns:
Tensor: [shape=y.shape].
"""
pi = 2 * torch.asin(torch.tensor(1.0))
x += ((x == 0) & (y == 0)) * 1.0
out = torch.atan(y / x)
out += ((y >= 0) & (x < 0)) * pi
out -= ((y < 0) & (x < 0)) * pi
out *= 1 - ((y > 0) & (x == 0)) * 1.0
out += ((y > 0) & (x == 0)) * (pi / 2)
out *= 1 - ((y < 0) & (x == 0)) * 1.0
out += ((y < 0) & (x == 0)) * (-pi / 2)
return out
# Define basic complex operations on torch.Tensor objects whose last dimension
# consists in the concatenation of the real and imaginary parts.
def _norm(x: torch.Tensor) -> torch.Tensor:
r"""Computes the norm value of a torch Tensor, assuming that it
comes as real and imaginary part in its last dimension.
Args:
x (Tensor): Input Tensor of shape [shape=(..., 2)]
Returns:
Tensor: shape as x excluding the last dimension.
"""
return torch.abs(x[..., 0]) ** 2 + torch.abs(x[..., 1]) ** 2
def _mul_add(a: torch.Tensor, b: torch.Tensor, out: Optional[torch.Tensor] = None) -> torch.Tensor:
"""Element-wise multiplication of two complex Tensors described
through their real and imaginary parts.
The result is added to the `out` tensor"""
# check `out` and allocate it if needed
target_shape = torch.Size([max(sa, sb) for (sa, sb) in zip(a.shape, b.shape)])
if out is None or out.shape != target_shape:
out = torch.zeros(target_shape, dtype=a.dtype, device=a.device)
if out is a:
real_a = a[..., 0]
out[..., 0] = out[..., 0] + (real_a * b[..., 0] - a[..., 1] * b[..., 1])
out[..., 1] = out[..., 1] + (real_a * b[..., 1] + a[..., 1] * b[..., 0])
else:
out[..., 0] = out[..., 0] + (a[..., 0] * b[..., 0] - a[..., 1] * b[..., 1])
out[..., 1] = out[..., 1] + (a[..., 0] * b[..., 1] + a[..., 1] * b[..., 0])
return out
def _mul(a: torch.Tensor, b: torch.Tensor, out: Optional[torch.Tensor] = None) -> torch.Tensor:
"""Element-wise multiplication of two complex Tensors described
through their real and imaginary parts
can work in place in case out is a only"""
target_shape = torch.Size([max(sa, sb) for (sa, sb) in zip(a.shape, b.shape)])
if out is None or out.shape != target_shape:
out = torch.zeros(target_shape, dtype=a.dtype, device=a.device)
if out is a:
real_a = a[..., 0]
out[..., 0] = real_a * b[..., 0] - a[..., 1] * b[..., 1]
out[..., 1] = real_a * b[..., 1] + a[..., 1] * b[..., 0]
else:
out[..., 0] = a[..., 0] * b[..., 0] - a[..., 1] * b[..., 1]
out[..., 1] = a[..., 0] * b[..., 1] + a[..., 1] * b[..., 0]
return out
def _inv(z: torch.Tensor, out: Optional[torch.Tensor] = None) -> torch.Tensor:
"""Element-wise multiplicative inverse of a Tensor with complex
entries described through their real and imaginary parts.
can work in place in case out is z"""
ez = _norm(z)
if out is None or out.shape != z.shape:
out = torch.zeros_like(z)
out[..., 0] = z[..., 0] / ez
out[..., 1] = -z[..., 1] / ez
return out
def _conj(z, out: Optional[torch.Tensor] = None) -> torch.Tensor:
"""Element-wise complex conjugate of a Tensor with complex entries
described through their real and imaginary parts.
can work in place in case out is z"""
if out is None or out.shape != z.shape:
out = torch.zeros_like(z)
out[..., 0] = z[..., 0]
out[..., 1] = -z[..., 1]
return out
def _invert(M: torch.Tensor, out: Optional[torch.Tensor] = None) -> torch.Tensor:
"""
Invert 1x1 or 2x2 matrices
Will generate errors if the matrices are singular: user must handle this
through his own regularization schemes.
Args:
M (Tensor): [shape=(..., nb_channels, nb_channels, 2)]
matrices to invert: must be square along dimensions -3 and -2
Returns:
invM (Tensor): [shape=M.shape]
inverses of M
"""
nb_channels = M.shape[-2]
if out is None or out.shape != M.shape:
out = torch.empty_like(M)
if nb_channels == 1:
# scalar case
out = _inv(M, out)
elif nb_channels == 2:
# two channels case: analytical expression
# first compute the determinent
det = _mul(M[..., 0, 0, :], M[..., 1, 1, :])
det = det - _mul(M[..., 0, 1, :], M[..., 1, 0, :])
# invert it
invDet = _inv(det)
# then fill out the matrix with the inverse
out[..., 0, 0, :] = _mul(invDet, M[..., 1, 1, :], out[..., 0, 0, :])
out[..., 1, 0, :] = _mul(-invDet, M[..., 1, 0, :], out[..., 1, 0, :])
out[..., 0, 1, :] = _mul(-invDet, M[..., 0, 1, :], out[..., 0, 1, :])
out[..., 1, 1, :] = _mul(invDet, M[..., 0, 0, :], out[..., 1, 1, :])
else:
raise Exception("Only 2 channels are supported for the torch version.")
return out
# Now define the signal-processing low-level functions used by the Separator
def expectation_maximization(
y: torch.Tensor,
x: torch.Tensor,
iterations: int = 2,
eps: float = 1e-10,
batch_size: int = 200,
):
r"""Expectation maximization algorithm, for refining source separation
estimates.
This algorithm allows to make source separation results better by
enforcing multichannel consistency for the estimates. This usually means
a better perceptual quality in terms of spatial artifacts.
The implementation follows the details presented in [1]_, taking
inspiration from the original EM algorithm proposed in [2]_ and its
weighted refinement proposed in [3]_, [4]_.
It works by iteratively:
* Re-estimate source parameters (power spectral densities and spatial
covariance matrices) through :func:`get_local_gaussian_model`.
* Separate again the mixture with the new parameters by first computing
the new modelled mixture covariance matrices with :func:`get_mix_model`,
prepare the Wiener filters through :func:`wiener_gain` and apply them
with :func:`apply_filter``.
References
----------
.. [1] S. Uhlich and M. Porcu and F. Giron and M. Enenkl and T. Kemp and
N. Takahashi and Y. Mitsufuji, `Improving music source separation based
on deep neural networks through data augmentation and network
blending.` 2017 IEEE International Conference on Acoustics, Speech
and Signal Processing (ICASSP). IEEE, 2017.
.. [2] N.Q. Duong and E. Vincent and R.Gribonval. `Under-determined
reverberant audio source separation using a full-rank spatial
covariance model.` IEEE Transactions on Audio, Speech, and Language
Processing 18.7 (2010): 1830-1840.
.. [3] A. Nugraha and A. Liutkus and E. Vincent. `Multichannel audio source
separation with deep neural networks.` IEEE/ACM Transactions on Audio,
Speech, and Language Processing 24.9 (2016): 1652-1664.
.. [4] A. Nugraha and A. Liutkus and E. Vincent. `Multichannel music
separation with deep neural networks.` 2016 24th European Signal
Processing Conference (EUSIPCO). IEEE, 2016.
.. [5] A. Liutkus and R. Badeau and G. Richard `Kernel additive models for
source separation.` IEEE Transactions on Signal Processing
62.16 (2014): 4298-4310.
Args:
y (Tensor): [shape=(nb_frames, nb_bins, nb_channels, 2, nb_sources)]
initial estimates for the sources
x (Tensor): [shape=(nb_frames, nb_bins, nb_channels, 2)]
complex STFT of the mixture signal
iterations (int): [scalar]
number of iterations for the EM algorithm.
eps (float or None): [scalar]
The epsilon value to use for regularization and filters.
Returns:
y (Tensor): [shape=(nb_frames, nb_bins, nb_channels, 2, nb_sources)]
estimated sources after iterations
v (Tensor): [shape=(nb_frames, nb_bins, nb_sources)]
estimated power spectral densities
R (Tensor): [shape=(nb_bins, nb_channels, nb_channels, 2, nb_sources)]
estimated spatial covariance matrices
Notes:
* You need an initial estimate for the sources to apply this
algorithm. This is precisely what the :func:`wiener` function does.
* This algorithm *is not* an implementation of the `exact` EM
proposed in [1]_. In particular, it does compute the posterior
covariance matrices the same (exact) way. Instead, it uses the
simplified approximate scheme initially proposed in [5]_ and further
refined in [3]_, [4]_, that boils down to just take the empirical
covariance of the recent source estimates, followed by a weighted
average for the update of the spatial covariance matrix. It has been
empirically demonstrated that this simplified algorithm is more
robust for music separation.
Warning:
It is *very* important to make sure `x.dtype` is `torch.float64`
if you want double precision, because this function will **not**
do such conversion for you from `torch.complex32`, in case you want the
smaller RAM usage on purpose.
It is usually always better in terms of quality to have double
precision, by e.g. calling :func:`expectation_maximization`
with ``x.to(torch.float64)``.
"""
# dimensions
(nb_frames, nb_bins, nb_channels) = x.shape[:-1]
nb_sources = y.shape[-1]
regularization = torch.cat(
(
torch.eye(nb_channels, dtype=x.dtype, device=x.device)[..., None],
torch.zeros((nb_channels, nb_channels, 1), dtype=x.dtype, device=x.device),
),
dim=2,
)
regularization = torch.sqrt(torch.as_tensor(eps)) * (
regularization[None, None, ...].expand((-1, nb_bins, -1, -1, -1))
)
# allocate the spatial covariance matrices
R = [torch.zeros((nb_bins, nb_channels, nb_channels, 2), dtype=x.dtype, device=x.device) for j in range(nb_sources)]
weight: torch.Tensor = torch.zeros((nb_bins,), dtype=x.dtype, device=x.device)
v: torch.Tensor = torch.zeros((nb_frames, nb_bins, nb_sources), dtype=x.dtype, device=x.device)
for it in range(iterations):
# constructing the mixture covariance matrix. Doing it with a loop
# to avoid storing anytime in RAM the whole 6D tensor
# update the PSD as the average spectrogram over channels
v = torch.mean(torch.abs(y[..., 0, :]) ** 2 + torch.abs(y[..., 1, :]) ** 2, dim=-2)
# update spatial covariance matrices (weighted update)
for j in range(nb_sources):
R[j] = torch.tensor(0.0, device=x.device)
weight = torch.tensor(eps, device=x.device)
pos: int = 0
batch_size = batch_size if batch_size else nb_frames
while pos < nb_frames:
t = torch.arange(pos, min(nb_frames, pos + batch_size))
pos = int(t[-1]) + 1
R[j] = R[j] + torch.sum(_covariance(y[t, ..., j]), dim=0)
weight = weight + torch.sum(v[t, ..., j], dim=0)
R[j] = R[j] / weight[..., None, None, None]
weight = torch.zeros_like(weight)
# cloning y if we track gradient, because we're going to update it
if y.requires_grad:
y = y.clone()
pos = 0
while pos < nb_frames:
t = torch.arange(pos, min(nb_frames, pos + batch_size))
pos = int(t[-1]) + 1
y[t, ...] = torch.tensor(0.0, device=x.device, dtype=x.dtype)
# compute mix covariance matrix
Cxx = regularization
for j in range(nb_sources):
Cxx = Cxx + (v[t, ..., j, None, None, None] * R[j][None, ...].clone())
# invert it
inv_Cxx = _invert(Cxx)
# separate the sources
for j in range(nb_sources):
# create a wiener gain for this source
gain = torch.zeros_like(inv_Cxx)
# computes multichannel Wiener gain as v_j R_j inv_Cxx
indices = torch.cartesian_prod(
torch.arange(nb_channels),
torch.arange(nb_channels),
torch.arange(nb_channels),
)
for index in indices:
gain[:, :, index[0], index[1], :] = _mul_add(
R[j][None, :, index[0], index[2], :].clone(),
inv_Cxx[:, :, index[2], index[1], :],
gain[:, :, index[0], index[1], :],
)
gain = gain * v[t, ..., None, None, None, j]
# apply it to the mixture
for i in range(nb_channels):
y[t, ..., j] = _mul_add(gain[..., i, :], x[t, ..., i, None, :], y[t, ..., j])
return y, v, R
def wiener(
targets_spectrograms: torch.Tensor,
mix_stft: torch.Tensor,
iterations: int = 1,
softmask: bool = False,
residual: bool = False,
scale_factor: float = 10.0,
eps: float = 1e-10,
):
"""Wiener-based separation for multichannel audio.
The method uses the (possibly multichannel) spectrograms of the
sources to separate the (complex) Short Term Fourier Transform of the
mix. Separation is done in a sequential way by:
* Getting an initial estimate. This can be done in two ways: either by
directly using the spectrograms with the mixture phase, or
by using a softmasking strategy. This initial phase is controlled
by the `softmask` flag.
* If required, adding an additional residual target as the mix minus
all targets.
* Refinining these initial estimates through a call to
:func:`expectation_maximization` if the number of iterations is nonzero.
This implementation also allows to specify the epsilon value used for
regularization. It is based on [1]_, [2]_, [3]_, [4]_.
References
----------
.. [1] S. Uhlich and M. Porcu and F. Giron and M. Enenkl and T. Kemp and
N. Takahashi and Y. Mitsufuji, `Improving music source separation based
on deep neural networks through data augmentation and network
blending.` 2017 IEEE International Conference on Acoustics, Speech
and Signal Processing (ICASSP). IEEE, 2017.
.. [2] A. Nugraha and A. Liutkus and E. Vincent. `Multichannel audio source
separation with deep neural networks.` IEEE/ACM Transactions on Audio,
Speech, and Language Processing 24.9 (2016): 1652-1664.
.. [3] A. Nugraha and A. Liutkus and E. Vincent. `Multichannel music
separation with deep neural networks.` 2016 24th European Signal
Processing Conference (EUSIPCO). IEEE, 2016.
.. [4] A. Liutkus and R. Badeau and G. Richard `Kernel additive models for
source separation.` IEEE Transactions on Signal Processing
62.16 (2014): 4298-4310.
Args:
targets_spectrograms (Tensor): spectrograms of the sources
[shape=(nb_frames, nb_bins, nb_channels, nb_sources)].
This is a nonnegative tensor that is
usually the output of the actual separation method of the user. The
spectrograms may be mono, but they need to be 4-dimensional in all
cases.
mix_stft (Tensor): [shape=(nb_frames, nb_bins, nb_channels, complex=2)]
STFT of the mixture signal.
iterations (int): [scalar]
number of iterations for the EM algorithm
softmask (bool): Describes how the initial estimates are obtained.
* if `False`, then the mixture phase will directly be used with the
spectrogram as initial estimates.
* if `True`, initial estimates are obtained by multiplying the
complex mix element-wise with the ratio of each target spectrogram
with the sum of them all. This strategy is better if the model are
not really good, and worse otherwise.
residual (bool): if `True`, an additional target is created, which is
equal to the mixture minus the other targets, before application of
expectation maximization
eps (float): Epsilon value to use for computing the separations.
This is used whenever division with a model energy is
performed, i.e. when softmasking and when iterating the EM.
It can be understood as the energy of the additional white noise
that is taken out when separating.
Returns:
Tensor: shape=(nb_frames, nb_bins, nb_channels, complex=2, nb_sources)
STFT of estimated sources
Notes:
* Be careful that you need *magnitude spectrogram estimates* for the
case `softmask==False`.
* `softmask=False` is recommended
* The epsilon value will have a huge impact on performance. If it's
large, only the parts of the signal with a significant energy will
be kept in the sources. This epsilon then directly controls the
energy of the reconstruction error.
Warning:
As in :func:`expectation_maximization`, we recommend converting the
mixture `x` to double precision `torch.float64` *before* calling
:func:`wiener`.
"""
if softmask:
# if we use softmask, we compute the ratio mask for all targets and
# multiply by the mix stft
y = (
mix_stft[..., None]
* (targets_spectrograms / (eps + torch.sum(targets_spectrograms, dim=-1, keepdim=True).to(mix_stft.dtype)))[
..., None, :
]
)
else:
# otherwise, we just multiply the targets spectrograms with mix phase
# we tacitly assume that we have magnitude estimates.
angle = atan2(mix_stft[..., 1], mix_stft[..., 0])[..., None]
nb_sources = targets_spectrograms.shape[-1]
y = torch.zeros(mix_stft.shape + (nb_sources,), dtype=mix_stft.dtype, device=mix_stft.device)
y[..., 0, :] = targets_spectrograms * torch.cos(angle)
y[..., 1, :] = targets_spectrograms * torch.sin(angle)
if residual:
# if required, adding an additional target as the mix minus
# available targets
y = torch.cat([y, mix_stft[..., None] - y.sum(dim=-1, keepdim=True)], dim=-1)
if iterations == 0:
return y
# we need to refine the estimates. Scales down the estimates for
# numerical stability
max_abs = torch.max(
torch.as_tensor(1.0, dtype=mix_stft.dtype, device=mix_stft.device),
torch.sqrt(_norm(mix_stft)).max() / scale_factor,
)
mix_stft = mix_stft / max_abs
y = y / max_abs
# call expectation maximization
y = expectation_maximization(y, mix_stft, iterations, eps=eps)[0]
# scale estimates up again
y = y * max_abs
return y
def _covariance(y_j):
"""
Compute the empirical covariance for a source.
Args:
y_j (Tensor): complex stft of the source.
[shape=(nb_frames, nb_bins, nb_channels, 2)].
Returns:
Cj (Tensor): [shape=(nb_frames, nb_bins, nb_channels, nb_channels, 2)]
just y_j * conj(y_j.T): empirical covariance for each TF bin.
"""
(nb_frames, nb_bins, nb_channels) = y_j.shape[:-1]
Cj = torch.zeros(
(nb_frames, nb_bins, nb_channels, nb_channels, 2),
dtype=y_j.dtype,
device=y_j.device,
)
indices = torch.cartesian_prod(torch.arange(nb_channels), torch.arange(nb_channels))
for index in indices:
Cj[:, :, index[0], index[1], :] = _mul_add(
y_j[:, :, index[0], :],
_conj(y_j[:, :, index[1], :]),
Cj[:, :, index[0], index[1], :],
)
return Cj
+237
View File
@@ -0,0 +1,237 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
import os
from typing import Dict
from seconohe.comfy_node_action import send_node_action
from seconohe.torch import get_torch_device_options, get_canonical_device
# ComfyUI imports
import folder_paths # ComfyUI's way to access model paths
# Local imports
# We are the main source, so we use the main_logger
from . import main_logger
from .utils.load_audio import audio_get_channels, force_stereo, force_sample_rate
from .db.models_db import ModelsDB
from .inference.demixer import get_demixer
DEF_MODEL = 'Kim_Vocal_2.safetensors'
DEF_ENTRY = 'Default'
MODELS_DIR = os.path.join(folder_paths.models_dir, "audio", "MDX")
DEMUCS_DIR = os.path.join(folder_paths.models_dir, "audio", "Demucs")
models_db_mdx = ModelsDB(MODELS_DIR)
models_db_demucs = ModelsDB(DEMUCS_DIR)
logger = main_logger
class AudioSeparateVocals:
PRIMARY_STEM = 'Vocals'
MODEL_T = 'MDX'
FILE_T = 'safetensors'
DEFAULT_MODEL = "Kim_Vocal_2.safetensors"
models_db = models_db_mdx
@classmethod
def _get_available_audio_models(cls):
# Refresh the database
cls.models_db.refresh()
# Filter the models this node can handle
cls.models_filtered = cls.models_db.get_filtered(primary_stem=cls.PRIMARY_STEM, model_t=cls.MODEL_T, file_t=cls.FILE_T,
default=cls.DEFAULT_MODEL, repeat_dl=True)
# We add any model downloaded and memorized by the GUI
return cls.models_filtered.get_display_names()
@classmethod
def INPUT_TYPES(cls):
device_options, default_device = get_torch_device_options()
return {
"required": {
"input_sound": ("AUDIO",),
"model": (cls._get_available_audio_models(),), # Dropdown for model selection
"segments": ("INT", {
"default": 1, # Default value
"min": 1, # Minimum allowed value
"max": 64, # Maximum allowed value (set a reasonable practical max)
"step": 1, # Step for slider/spinbox
"display": "slider" # How to display: "number" or "slider"
}),
"target_device": (device_options, {
"default": default_device,
"tooltip": "The device (CPU or CUDA) to which the projection layer will be assigned for computation."}),
}
}
RETURN_TYPES = ("AUDIO", "AUDIO",)
RETURN_NAMES = (PRIMARY_STEM, "Complement",)
FUNCTION = "execute"
CATEGORY = "audio/separation"
DESCRIPTION = "Separates vocals using MDX-Net networks"
UNIQUE_NAME = "AudioSeparateVocals"
DISPLAY_NAME = "Vocals using MDX"
def __init__(self):
super().__init__()
self.demixer = None
def execute(self, input_sound: Dict, model: str, segments: int, target_device: str):
# Get information for the selected model
main_logger.info(f"Selected model: {model}")
model_data = self.models_filtered.get_by_display_name(model)
if model_data is None:
raise ValueError("Unknown model selected, please refresh pressing `R` and select another")
model_path = model_data.get('model_path')
# Create or recycle a demixer
device = get_canonical_device(target_device)
if self.demixer is None or self.demixer.d['hash'] != model_data['hash']:
# New demixer
logger.debug("Creating a new demixer object")
# This will load the model, optionally downloading it
self.demixer = get_demixer(model_data, device, MODELS_DIR)
else:
# Update the device
self.demixer.set_device(device)
# Handle a change in the icon of the model name
if model_path is None:
# Was downloaded
send_node_action(logger, "change_widget", "model", model_data['indicator'] + model_data['filtered_name'])
# Match channels and S/R
waveform = input_sound['waveform']
sample_rate = input_sound['sample_rate']
if audio_get_channels(waveform) == 1 and self.demixer.ch == 2:
waveform = force_stereo(waveform)
if sample_rate != self.demixer.sr:
waveform = force_sample_rate(waveform, sample_rate, self.demixer.sr)
# Demix
wavs = self.demixer(waveform, segments)
return (wavs[0], wavs[1],)
class AudioSeparateInstrumental(AudioSeparateVocals):
PRIMARY_STEM = 'Instrumental'
DEFAULT_MODEL = "Kim_Inst.safetensors"
DESCRIPTION = "Separates instruments using MDX-Net networks"
UNIQUE_NAME = "AudioSeparateInstrumental"
DISPLAY_NAME = "Instrumental using MDX"
RETURN_NAMES = (PRIMARY_STEM, "Complement",)
class AudioSeparateBass(AudioSeparateVocals):
PRIMARY_STEM = 'Bass'
DEFAULT_MODEL = "kuielab_b_bass.safetensors"
DESCRIPTION = "Separates bass using MDX-Net networks"
UNIQUE_NAME = "AudioSeparateBass"
DISPLAY_NAME = "Bass using MDX"
RETURN_NAMES = (PRIMARY_STEM, "Complement",)
class AudioSeparateDrums(AudioSeparateVocals):
PRIMARY_STEM = 'Drums'
DEFAULT_MODEL = "kuielab_b_drums.safetensors"
DESCRIPTION = "Separates drums using MDX-Net networks"
UNIQUE_NAME = "AudioSeparateDrums"
DISPLAY_NAME = "Drums using MDX"
RETURN_NAMES = (PRIMARY_STEM, "Complement",)
class AudioSeparateVarious(AudioSeparateVocals):
PRIMARY_STEM = ["Other", "Reverb"]
DEFAULT_MODEL = "Reverb_HQ_By_FoxJoy.safetensors"
DESCRIPTION = "Misc. separators using MDX-Net networks"
UNIQUE_NAME = "AudioSeparateVarious"
DISPLAY_NAME = "Various using MDX"
RETURN_NAMES = ("Main", "Complement",)
class AudioSeparateDemucs(AudioSeparateVocals):
PRIMARY_STEM = None
MODEL_T = 'Demucs'
FILE_T = 'safetensors'
DEFAULT_MODEL = "htdemucs_ft.safetensors"
models_db = models_db_demucs
@classmethod
def INPUT_TYPES(cls):
device_options, default_device = get_torch_device_options()
return {
"required": {
"input_sound": ("AUDIO",),
"model": (cls._get_available_audio_models(),), # Dropdown for model selection
"shifts": ("INT", {
"default": 0, "min": 0, "max": 16, "step": 1, "display": "slider",
"tooltip": "Number of random shifts for equivariant stabilization.\n"
"Higher values improve quality but are slower. 0 disables it."
}),
"overlap": ("FLOAT", {
"default": 0.25, "min": 0.0, "max": 0.99, "step": 0.01, "display": "slider",
"tooltip": "Amount of overlap between audio chunks.\n"
"Higher values can reduce stitching artifacts but are slower."
}),
"custom_segment": ("BOOLEAN", {
"default": False,
"label_on": "enabled",
"label_off": "disabled",
"tooltip": "Enable to override the model's default segment length.\n"
"Disabling uses the recommended length from the model file.\n"
"Useful for HDemucs and Demucs models, not much for HTDemucs."
}),
"segment": ("INT", {
"default": 44, "min": 10, "max": 120, "step": 1, "display": "slider",
"tooltip": "Length of audio chunks to process at a time (in seconds).\n"
"Higher values need more VRAM but can improve quality."
}),
"target_device": (device_options, {
"default": default_device,
"tooltip": "The device (CPU or CUDA) to which the projection layer will be assigned for computation."}),
}
}
# Define the output names. We define all 6 possible stems.
# The execution logic will handle returning 'None' for unused outputs.
RETURN_NAMES = ("vocals", "drums", "bass", "other", "guitar", "piano")
RETURN_TYPES = ("AUDIO", "AUDIO", "AUDIO", "AUDIO", "AUDIO", "AUDIO")
DESCRIPTION = "Demucs Audio Separator (4/6 stems)"
UNIQUE_NAME = "AudioSeparateDemucs"
DISPLAY_NAME = "Demucs Audio Separator"
def execute(self, input_sound: Dict, model: str, shifts: int, overlap: float, custom_segment: bool, segment: int,
target_device: str):
# Get information for the selected model
main_logger.info(f"Selected model: {model}")
model_data = self.models_filtered.get_by_display_name(model)
if model_data is None:
raise ValueError("Unknown model selected, please refresh pressing `R` and select another")
model_path = model_data.get('model_path')
# Create or recycle a demixer
device = get_canonical_device(target_device)
if self.demixer is None or self.demixer.d['hash'] != model_data['hash']:
# New demixer
logger.debug("Creating a new demixer object")
# This will load the model, optionally downloading it
self.demixer = get_demixer(model_data, device, DEMUCS_DIR)
else:
# Update the device
self.demixer.set_device(device)
# Handle a change in the icon of the model name
if model_path is None:
# Was downloaded
send_node_action(logger, "change_widget", "model", model_data['indicator'] + model_data['filtered_name'])
# Match channels and S/R
waveform = input_sound['waveform']
sample_rate = input_sound['sample_rate']
if audio_get_channels(waveform) == 1 and self.demixer.ch == 2:
waveform = force_stereo(waveform)
if sample_rate != self.demixer.sr:
waveform = force_sample_rate(waveform, sample_rate, self.demixer.sr)
# Demix
wavs = self.demixer(waveform, shifts=shifts, overlap=overlap, segment=segment if custom_segment else None)
return tuple(wavs)
@@ -8,7 +8,7 @@
import logging
import torch
import torchaudio
from .misc import NODES_NAME
from .. import NODES_NAME
logger = logging.getLogger(f"{NODES_NAME}.load_audio")
@@ -8,7 +8,7 @@ import importlib
import logging
import os
import sys
from .misc import NODES_NAME
from .. import NODES_NAME
logger = logging.getLogger(f"{NODES_NAME}.load_class")
@@ -11,7 +11,7 @@ try:
with_onnx = True
except Exception:
with_onnx = False
from .misc import NODES_NAME
from .. import NODES_NAME
logger = logging.getLogger(f"{NODES_NAME}.load_onnx")
+62
View File
@@ -0,0 +1,62 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# Model loader helper
# Original code from Gemini 2.5 Pro
import logging
from safetensors.torch import load_file
from .. import NODES_NAME
logger = logging.getLogger(f"{NODES_NAME}.load_safetensors")
def load_state_dict(model_run, state_dict):
try:
missing_keys, unexpected_keys = model_run.load_state_dict(state_dict, strict=False)
if missing_keys:
logger.warning(f"Missing keys in state_dict for model_run: {missing_keys}")
if unexpected_keys:
logger.warning(f"Unexpected keys in state_dict for model_run: {unexpected_keys}")
if not missing_keys and not unexpected_keys:
logger.debug("All keys matched successfully.")
except RuntimeError as e:
logger.error(f"RuntimeError during model_run.load_state_dict: {e}")
logger.error("This might indicate a mismatch between saved weights and model architecture.")
raise
def load_safetensors(model_path, model_run, device):
logger.info("Loading PyTorch model from .safetensors file...")
# 1. Load the state_dict from the file, EXPLICITLY forcing all tensors onto the CPU.
state_dict = load_file(model_path, device="cpu")
# 2. Load the CPU state_dict into the CPU model. This is now a safe operation.
if hasattr(model_run, 'signatures'):
# This is the way we store Demucs models
signatures = model_run.signatures
if len(signatures) == 1:
# Single model
prefix = f"{signatures[0]}."
logger.debug(f"Single model with filtered keys, prefix: {prefix}")
# Just remove the prefix
# Note: child models needs the if k.startswith(prefix) because they contain extra keys
state_dict = {k[len(prefix):]: v for k, v in state_dict.items() if k.startswith(prefix)}
# The rest is as a regular model
else:
# Bag of models, load the keys for each sub-model
logger.debug("Multiple models with filtered keys")
# Load weights, but just once when the same model is used more than once
for sig, sub_model in {s: m for m, s in zip(model_run.models, signatures)}.items():
prefix = f"{sig}."
logger.debug(f" - Filtering {prefix}")
sub_state_dict = {k[len(prefix):]: v for k, v in state_dict.items() if k.startswith(prefix)}
load_state_dict(sub_model, sub_state_dict)
# Finished
return model_run
load_state_dict(model_run, state_dict)
model_run.target_device = device
return model_run
+51
View File
@@ -0,0 +1,51 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
import argparse
from fractions import Fraction
import json
from .. import __version__, __copyright__, __license__, __author__
def cli_add_verbose(parser):
parser.add_argument('-v', '--verbose', action='count', default=0,
help="Enable verbose output to see details of the process.")
class PrintVersionAction(argparse.Action):
def __init__(self, option_strings, dest, nargs=None, **kwargs):
super().__init__(option_strings, dest, nargs=0, **kwargs)
def __call__(self, parser, namespace, values, option_string=None):
# Format the version information
version_info = f"""{parser.prog} (Audio Separation) {__version__}
{__copyright__}
{__license__}
This is free software: you are free to change and redistribute it.
There is NO WARRANTY, to the extent permitted by law.
Written by {__author__}"""
print(version_info)
# Exit the parser
parser.exit()
def cli_add_version(parser, prog_name):
parser.add_argument('-V', '--version', help="Show version and copyright information and exit",
action=PrintVersionAction)
class FractionEncoder(json.JSONEncoder):
def default(self, obj):
if isinstance(obj, Fraction):
# Represent the Fraction as a dictionary with a type hint
return {'_type': 'Fraction', 'numerator': obj.numerator, 'denominator': obj.denominator}
return super().default(obj)
def json_object_hook(d):
"""The decoder hook for our custom Fraction serialization."""
if d.get('_type') == 'Fraction':
return Fraction(d['numerator'], d['denominator'])
return d
@@ -8,7 +8,7 @@
import logging
import os
import torchaudio
from .misc import NODES_NAME
from .. import NODES_NAME
logger = logging.getLogger(f"{NODES_NAME}.save_audio")
+14
View File
@@ -13,6 +13,13 @@ Might work for other ONNX files using the same operands.
But you'll need a PyTorch class for its architecture that creates the layers
in the same order as the ONNX file.
# Demucs to safetensors
File: demucs2safetensors.py
From a YAML file and the PyTorch components it creates safetensors version of the Demucs file.
All metadata and models are contained in the new file.
# Batch converter
File: batch_convert.py
@@ -59,6 +66,13 @@ Used to display our PyTorch class, also the state_dict keys.
It can optionally export the class structure as an ONNX file that can be loaded by
Netron, but contains too much extra names.
# Show Metadata
File: show_metadata.py
Used to display our safetensors metadata.
It can decode structures in JSON format we use for Demucs.
# Style Fixer
File: style_fixer.py
Regular → Executable
+8 -5
View File
@@ -1,3 +1,4 @@
#!/usr/bin/env python3
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
@@ -11,14 +12,15 @@ import argparse
import json
import os
import re
from seconohe.logger import logger_set_standalone
import subprocess
import sys
# Local imports
import bootstrap # noqa: F401
from source.db.hash import get_hash
from source.db.models_db import load_known_models, save_known_models, get_db_filename
from source.utils.logger import main_logger, logger_set_standalone
from source.utils.misc import cli_add_verbose
from src.nodes import main_logger
from src.nodes.db.hash import get_hash
from src.nodes.db.models_db import load_known_models, save_known_models, get_db_filename
from src.nodes.utils.misc import cli_add_verbose, cli_add_version
def parse_converter_output(output):
@@ -37,7 +39,7 @@ def parse_converter_output(output):
def main(args):
logger_set_standalone(args)
logger_set_standalone(main_logger, args)
# 1. Load the JSON metadata file
model_db = load_known_models(args.json_file)
if model_db is None:
@@ -200,6 +202,7 @@ if __name__ == "__main__":
parser.add_argument('--model_location', type=str, default='source/inference/MDX_Net.py:MDX_Net',
help="Python path to the model class.")
cli_add_verbose(parser)
cli_add_version(parser, __name__)
args = parser.parse_args()
main(args)
+80
View File
@@ -0,0 +1,80 @@
#!/usr/bin/env python3
import re
import sys
from pathlib import Path
# --- Configuration ---
# The path to your pyproject.toml file, relative to the project root
PYPROJECT_PATH = Path("pyproject.toml")
# The path to the Python file containing the __version__ string
SOURCE_VERSION_PATH = Path("src/nodes/__init__.py")
# --- End Configuration ---
def get_version_from_pyproject(file_path: Path) -> str | None:
"""Extracts the version string from a pyproject.toml file."""
try:
content = file_path.read_text()
# A simple regex to find `version = "..."` under the `[project]` table
match = re.search(r'^version\s*=\s*"(.*?)"', content, re.MULTILINE)
if match:
return match.group(1)
except FileNotFoundError:
print(f"Error: {file_path} not found.", file=sys.stderr)
except Exception as e:
print(f"Error reading or parsing {file_path}: {e}", file=sys.stderr)
return None
def get_version_from_source(file_path: Path) -> str | None:
"""Extracts the __version__ string from a Python source file."""
try:
content = file_path.read_text()
# A simple regex to find `__version__ = "..."`
match = re.search(r'^__version__\s*=\s*"(.*?)"', content, re.MULTILINE)
if match:
return match.group(1)
except FileNotFoundError:
print(f"Error: {file_path} not found.", file=sys.stderr)
except Exception as e:
print(f"Error reading or parsing {file_path}: {e}", file=sys.stderr)
return None
def main() -> int:
"""
Compares version strings from pyproject.toml and the source code.
Exits with a non-zero status code if they do not match.
"""
print("--- Checking version consistency ---")
# Get versions
pyproject_version = get_version_from_pyproject(PYPROJECT_PATH)
source_version = get_version_from_source(SOURCE_VERSION_PATH)
# Validate that we found both
if not pyproject_version:
print(f"Error: Could not find version in {PYPROJECT_PATH}", file=sys.stderr)
return 1
if not source_version:
print(f"Error: Could not find `__version__` in {SOURCE_VERSION_PATH}", file=sys.stderr)
return 1
print(f"Version in {PYPROJECT_PATH}: {pyproject_version}")
print(f"Version in {SOURCE_VERSION_PATH}: {source_version}")
# Compare and exit
if pyproject_version == source_version:
print("✅ Versions are consistent.")
return 0
else:
print("\n❌ Error: Version mismatch!", file=sys.stderr)
print(f" pyproject.toml has version '{pyproject_version}'", file=sys.stderr)
print(f" {SOURCE_VERSION_PATH} has version '{source_version}'", file=sys.stderr)
print(" Please ensure both versions are identical.", file=sys.stderr)
return 1
if __name__ == "__main__":
sys.exit(main())
Regular → Executable
+20 -14
View File
@@ -1,3 +1,4 @@
#!/usr/bin/env python3
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
@@ -7,16 +8,17 @@
# Run it using: python tool/demix.py -m HASH AUDIO
import argparse
import os
from seconohe.logger import logger_set_standalone
from seconohe.torch import get_torch_device_options, get_canonical_device
import sys
import torch
# Local imports
import bootstrap # noqa: F401
from source.db.models_db import ModelsDB, cli_add_models_and_db
from source.inference.demixer import get_demixer
from source.utils.logger import main_logger, logger_set_standalone
from source.utils.load_audio import load_audio
from source.utils.save_audio import save_audio
from source.utils.misc import cli_add_verbose
from src.nodes import main_logger
from src.nodes.db.models_db import ModelsDB, cli_add_models_and_db
from src.nodes.inference.demixer import get_demixer
from src.nodes.utils.load_audio import load_audio
from src.nodes.utils.misc import cli_add_verbose, cli_add_version
from src.nodes.utils.save_audio import save_audio
BANNER = "🎵 MDX-Net Audio Separation Tool 🎵"
@@ -24,7 +26,8 @@ BANNER = "🎵 MDX-Net Audio Separation Tool 🎵"
# --- Main Demixing Logic ---
def demix(d, args):
main_logger.info("🚀 Starting audio separation process...")
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
_, device = get_torch_device_options()
device = get_canonical_device(device)
main_logger.info(f"💻 Using device: {device}")
# --- Load and Prepare Audio ---
@@ -34,7 +37,7 @@ def demix(d, args):
demixer = get_demixer(d, device, args.models_dir)
# --- Do inference in chunks ---
wavs = demixer(waveform, args.segments)
wavs = demixer(waveform, args.segments if args.segments else None)
# --- Save outputs ---
base, ext = os.path.splitext(args.input_file)
@@ -46,6 +49,8 @@ def demix(d, args):
if args.save_complement:
for wav in wavs[1:]:
if wav is None or not wav['generated']:
continue
out_path = f"{args.out_base or base}_{wav['stem']}{out_ext}"
save_audio(wav['waveform'], wav['sample_rate'], out_path, out_format)
@@ -70,21 +75,22 @@ if __name__ == "__main__":
parser.add_argument('--no_main', dest='save_main', action='store_false',
help="Do not save the main separated stem.")
parser.add_argument('--save_complement', action='store_true',
help="Save the complement stem (input - main).")
help="Save all the stems, including the complement stem (input - main).")
parser.add_argument('--out_base', type=str, default=None,
help="Base for the output path. No extension here, we will add the name of the stem and extension")
parser.add_argument('--format', type=str, default=None, choices=['wav', 'flac', 'mp3'],
help="Output audio format. Defaults to input format.")
# --- Control Arguments ---
parser.add_argument('--segments', type=int, default=1,
help="How many audio segments to process at once")
parser.add_argument('--segments', type=int, default=0,
help="How many audio segments to process at once. 0 uses the value from the model")
parser.add_argument('-l', '--list', action='store_true', help="Show available models")
cli_add_verbose(parser)
cli_add_version(parser, __name__)
parser.set_defaults(save_main=True)
args = parser.parse_args()
logger_set_standalone(args)
logger_set_standalone(main_logger, args)
main_logger.info(BANNER)
# Sanity check
@@ -117,7 +123,7 @@ if __name__ == "__main__":
# Look for the selected model
d = models.get(args.model)
if d is None:
main_logger.error(f"💥 Unknown model `{args.model}`.")
main_logger.error(f"💥 Unknown model `{args.model}`. If you provided a file check it exists")
sys.exit(3)
try:
main_logger.info(f"📂 Using model from `{d['model_path']}`")
+290
View File
@@ -0,0 +1,290 @@
#!/usr/bin/env python3
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# Tool to convert a Demucs v3/4 model into safetensors
# python tool/tool/demucs2safetensors.py --yaml DEMUCS.YAML
# You must manually download the .th files and copy them to the YAML dir
# First version by Gemini 2.5 Pro
from contextlib import contextmanager
from copy import deepcopy
import argparse
import json
from pathlib import Path
from safetensors.torch import save_file
from seconohe.logger import debugl, logger_set_standalone
import sys
import torch
from typing import Dict
import yaml
try:
# We need the original Demucs library to dequantize
from demucs.states import set_state # noqa: F401
with_demuc_lib = True
except Exception:
with_demuc_lib = False
import bootstrap # noqa: F401
from src.nodes import main_logger
from src.nodes.utils.misc import cli_add_verbose, FractionEncoder, cli_add_version
from src.nodes.db.models_db import cli_add_db, get_download_url, ModelsDB
from src.nodes.db.hash import get_hash
import src.nodes.inference.Demucs as local_demucs_module
import src.nodes.inference.HDemucs as local_hdemucs_module
import src.nodes.inference.HTDemucs as local_htdemucs_module
MODULES_MAP = {'demucs': local_demucs_module,
'demucs.demucs': local_demucs_module,
'demucs.hdemucs': local_hdemucs_module,
'demucs.htdemucs': local_htdemucs_module}
MAP = {'freq_encoder': 'encoder',
'freq_decoder': 'decoder',
'time_encoder': 'tencoder',
'time_decoder': 'tdecoder'}
logger = main_logger
@contextmanager
def remap_module(modules_map):
"""
A context manager to temporarily remap an old module name to a new one.
This is useful for loading pickled objects that depend on old paths.
"""
original_modules = {}
for old_name, new_module in modules_map.items():
original_modules[old_name] = sys.modules.get(old_name)
sys.modules[old_name] = new_module
try:
yield
finally:
# Restore the original state
for old_name, original_module in original_modules.items():
if original_module is not None:
sys.modules[old_name] = original_module
else:
# If the module wasn't there before, remove our patch
del sys.modules[old_name]
def solve_simple_pt(yaml_data, pkg):
""" PyTorch Audio lib has some raw models, we store the metadata in the YAML """
klass = yaml_data.get('klass')
if klass is None:
main_logger.error("No `klass` in YAML")
sys.exit(4)
if klass == 'Demucs':
klass = local_demucs_module.Demucs
elif klass == 'HDemucs':
klass = local_hdemucs_module.HDemucs
elif klass == 'HTDemucs':
klass = local_htdemucs_module.HTDemucs
else:
main_logger.error("Unknown model `klass` {klass}")
sys.exit(4)
args = yaml_data.get('args', {})
kwargs = yaml_data.get('kwargs', {})
# For PyTorch Audio model (very old code?)
new_dict = {}
for k, v in pkg.items():
parts = k.split('.')
gr = parts[0]
if gr in MAP:
k = MAP[gr] + '.' + '.'.join(parts[1:])
new_dict[k] = v
return klass, args, kwargs, new_dict
def convert_demucs_model(yaml_path_str: str, model_paths: list[str], output_path_str: str, all_metadata, data: Dict):
"""
Loads an original Demucs model bag, extracts weights and all necessary metadata
from the YAML and .th files, and saves it to a single, secure .safetensors file.
"""
yaml_path = Path(yaml_path_str)
output_path = Path(output_path_str)
# 1. Load and parse the YAML file
main_logger.info(f"\n--- 🚀 Starting conversion for {yaml_path.name} ---\n")
main_logger.info(f"- Loading YAML definition from: {yaml_path}")
with open(yaml_path, 'r') as f:
yaml_data = yaml.safe_load(f)
signatures = yaml_data['models']
main_logger.info(f"- Found {len(signatures)} model signatures in YAML: {signatures}")
# 2. Match the provided .th files to their signatures
model_file_map = {}
for path_str in model_paths:
path = Path(path_str)
# The signature is the part of the filename before the first '-' or '.'
sig = path.stem.split('-')[0]
if sig not in signatures:
main_logger.warning(f"File {path.name} with signature {sig} is not listed in the YAML file. Skipping.")
continue
if sig in model_file_map:
raise ValueError(f"Duplicate files found for signature {sig}.")
model_file_map[sig] = path
# Verify that all signatures from the YAML have a corresponding file
if len(model_file_map) != len(set(signatures)):
missing_sigs = set(signatures) - set(model_file_map.keys())
raise FileNotFoundError(f"Missing model files for signatures: {missing_sigs}")
main_logger.info("- Successfully mapped all signatures to model files.")
# 3. Load model packages and extract metadata and state dicts
# all_metadata = {}
consolidated_state_dict = {}
for sig in signatures:
path = model_file_map[sig]
main_logger.info(f"\n- Processing model '{sig}' from '{path.name}'...")
# Use the context manager to perform the remap
with remap_module(MODULES_MAP):
# This is the only "unsafe" part, loading the original pickle file
pkg = torch.load(path, map_location='cpu', weights_only=False)
debugl(logger, 2, f"PyTorch data type is {type(pkg)}")
if isinstance(pkg, dict) and 'klass' in pkg:
klass, args, kwargs, state = pkg["klass"], pkg["args"], pkg["kwargs"], pkg["state"]
else:
klass, args, kwargs, state = solve_simple_pt(yaml_data, pkg)
main_logger.info(f" - Model class: {klass.__module__}.{klass.__name__}")
# Dequantize if necessary by letting the original code handle it
if state.get('__quantized'):
if not with_demuc_lib:
main_logger.error("Don't use quantized models. Look for the same model without `_q`")
main_logger.error("Alternatively install the demucs Python module")
sys.exit(3)
else:
main_logger.info(" - Model is quantized. Dequantizing weights...")
model_instance = klass(*args, **kwargs)
set_state(model_instance, state)
clean_state_dict = model_instance.state_dict()
else:
clean_state_dict = state
# Store this model's metadata, keyed by its signature
all_metadata[sig] = json.dumps({
'class_module': klass.__module__,
'class_name': klass.__name__,
'args': args,
'kwargs': kwargs,
}, cls=FractionEncoder) # Use the custom encoder here
# Add the weights to the consolidated dict, prefixed by signature
for key, value in clean_state_dict.items():
consolidated_state_dict[f"{sig}.{key}"] = value
# 4. Add the top-level YAML data to the metadata
all_metadata['is_bag_of_models'] = json.dumps(len(set(signatures)) > 1)
all_metadata['signatures'] = json.dumps(signatures)
if 'weights' in yaml_data:
all_metadata['weights'] = json.dumps(yaml_data['weights'])
if 'segment' in yaml_data:
all_metadata['segment'] = str(yaml_data['segment'])
# 5. Calculate the total number of parameters from the final state dict
total_params = sum(p.numel() for p in consolidated_state_dict.values())
main_logger.info(f"- Total number of parameters in the model: {total_params:,}")
# Add the count as a string to the metadata dictionary
data['params'] = all_metadata['params'] = str(total_params)
# 6. Save the final .safetensors file
main_logger.info(f"- \U0001F4BE Saving consolidated model and metadata to: {output_path}")
save_file(consolidated_state_dict, output_path, metadata=all_metadata)
main_logger.info("\n--- 🎉 Conversion complete! ---")
if __name__ == '__main__':
parser = argparse.ArgumentParser(description="Convert Demucs .th models to a single .safetensors file.")
parser.add_argument('--yaml', required=True, type=str, help="Path to the Demucs .yaml file.")
parser.add_argument('--models', nargs='*', default=None, type=str,
help="Optional. Paths to .th files. If not provided, assumes they are "
"in the same directory as the YAML.")
parser.add_argument('--output', default=None, type=str,
help="Optional. Path for the output .safetensors file. If not provided, "
"it's saved next to the YAML with the same name.")
cli_add_db(parser)
cli_add_verbose(parser)
cli_add_version(parser, __name__)
args = parser.parse_args()
logger_set_standalone(main_logger, args)
main_logger.info("⚙️ PyTorch Demucs to Safetensors converter\n")
# Get information about the YAML file in our database
yaml_path = Path(args.yaml)
db = ModelsDB(yaml_path.parent)
known = db.get_filtered(model_t="Demucs")
hash = get_hash(yaml_path)
d = known.get_by_hash(hash)
if d is None:
# We don't have the hash for it
d = known.get_by_file_name(yaml_path.name)
if d is None:
# We use some information from the database to populate the metadata so is better if we have the
# information in the database
main_logger.error(f"{yaml_path} not in data base, please add it first")
sys.exit(3)
else:
main_logger.error(f"{yaml_path} in database, but with different hash")
logger.info(f"Model description: {d['desc']}")
# Logic for optional --models
model_files = args.models
if not model_files:
main_logger.debug("No --models provided. Searching for .th files alongside the YAML...")
yaml_dir = yaml_path.parent
with open(yaml_path, 'r') as f:
signatures = yaml.safe_load(f)['models']
model_files = []
for sig in set(signatures): # Use set to avoid redundant searches
found = list(yaml_dir.glob(f'{sig}*.th')) + list(yaml_dir.glob(f'{sig}*.pt'))
if not found:
raise FileNotFoundError(f"Could not automatically find a model file for signature '{sig}' in {yaml_dir}")
if len(found) > 1:
full_match = yaml_dir / (sig + '.th') # UVR Demucs uses it
if full_match in found:
found = [full_match]
else:
main_logger.warning(f"Found multiple files for signature '{sig}', using the first one: {found[0]}")
model_files.append(str(found[0]))
main_logger.info(f"Automatically found model files: {model_files}")
# Logic for optional --output
output_file = args.output
if not output_file:
output_file = yaml_path.with_suffix('.safetensors')
main_logger.info(f"Defaulting to: `{output_file}` (No --output provided)")
else:
output_file = Path(output_file)
# Adjust the data to the converted version
metadata = {}
d = deepcopy(d)
d["download"] = "Main/Demucs"
d["name"] = Path(d["name"]).stem + ".safetensors"
d["file_t"] = "safetensors"
# Add it to the safetensors
metadata["desc"] = d["desc"]
metadata["download"] = get_download_url(d)
metadata["file_t"] = "safetensors"
metadata["model_t"] = "Demucs"
metadata["name"] = d["name"]
metadata["primary_stem"] = json.dumps(d["primary_stem"])
metadata["project"] = "https://github.com/set-soft/AudioSeparation"
convert_demucs_model(args.yaml, model_files, output_file, metadata, d)
# Now update the DB
db.remove(known.get_by_file_name(d["name"]))
db.add(get_hash(output_file), d)
db.save()
Regular → Executable
+1
View File
@@ -1,3 +1,4 @@
#!/usr/bin/env python3
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
Regular → Executable
+9 -6
View File
@@ -1,3 +1,4 @@
#!/usr/bin/env python3
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
@@ -12,16 +13,17 @@ import json
import numpy as np
import onnx
import os
from safetensors.torch import save_file
from seconohe.logger import logger_set_standalone
import sys
import torch
from torch import nn
# Local imports
import bootstrap # noqa: F401
from safetensors.torch import save_file
from source.utils.load_class import import_model_class
from source.utils.logger import main_logger, logger_set_standalone
from source.utils.misc import cli_add_verbose
from source.db.models_db import get_download_url
from src.nodes import main_logger
from src.nodes.utils.load_class import import_model_class
from src.nodes.utils.misc import cli_add_verbose, cli_add_version
from src.nodes.db.models_db import get_download_url
class OnnxGraph:
@@ -224,9 +226,10 @@ if __name__ == "__main__":
"(default: source/inference/MDX_Net.py:MDX_Net)")
parser.add_argument('-j', '--metadata', type=str, help="Metadata to include in the hyperparameters, JSON format")
cli_add_verbose(parser)
cli_add_version(parser, __name__)
args = parser.parse_args()
logger_set_standalone(args)
logger_set_standalone(main_logger, args)
if args.metadata is not None:
try:
args.metadata = json.loads(args.metadata)
Regular → Executable
+8 -5
View File
@@ -1,3 +1,4 @@
#!/usr/bin/env python3
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
@@ -6,13 +7,14 @@
# Tool to show PyTorch class i.e:
# python tool/show_class.py -m source/inference/MDX_Net.py:MDX_Net
import argparse
from seconohe.logger import logger_set_standalone
import sys
from torch import nn
# Local imports
import bootstrap # noqa: F401
from source.utils.logger import main_logger, logger_set_standalone
from source.utils.load_class import import_model_class
from source.utils.misc import cli_add_verbose
from src.nodes import main_logger
from src.nodes.utils.load_class import import_model_class
from src.nodes.utils.misc import cli_add_verbose, cli_add_version
def load_class(args):
@@ -121,16 +123,17 @@ if __name__ == "__main__":
parser.add_argument('-n', '--num_stages', type=int, default=5,
choices=[2, 3, 4, 5, 6, 7], # Restrict to known valid values
help="The number of U-Net stages in the model. (choices: 2 to 7, default: 5)")
cli_add_verbose(parser)
parser.add_argument('-o', '--export_onnx', type=str, default=None,
help="Path for the optional output .onnx file.\n"
"Only the structure is exported")
parser.add_argument('-k', '--keys', action='store_true', help="Print the keys for the state_dict.")
parser.add_argument('-C', '--compact', action='store_true', help="Print a compact representation.")
parser.add_argument('-S', '--no_show', action='store_false', help="Don't print the structure.")
cli_add_verbose(parser)
cli_add_version(parser, __name__)
args = parser.parse_args()
logger_set_standalone(args)
logger_set_standalone(main_logger, args)
model = load_class(args)
if args.no_show:
show(model)
Regular → Executable
+8 -5
View File
@@ -1,3 +1,4 @@
#!/usr/bin/env python3
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
@@ -8,13 +9,14 @@
# Run it using: python tool/show_db.py
import argparse
import pprint
from seconohe.logger import logger_set_standalone
import sys
# Local imports
import bootstrap # noqa: F401
from source.db.hash_dir import hash_dir
from source.db.models_db import load_known_models, cli_add_models_and_db, save_known_models, get_models
from source.utils.logger import main_logger, logger_set_standalone
from source.utils.misc import cli_add_verbose
from src.nodes import main_logger
from src.nodes.db.hash_dir import hash_dir
from src.nodes.db.models_db import load_known_models, cli_add_models_and_db, save_known_models, get_models
from src.nodes.utils.misc import cli_add_verbose, cli_add_version
# Do nothing, you can apply some change here
@@ -63,7 +65,7 @@ def apply_process(model_db):
def main(args):
logger_set_standalone(args)
logger_set_standalone(main_logger, args)
# 1. Load the JSON metadata file
model_db = load_known_models(args.json_file)
if model_db is None:
@@ -128,6 +130,7 @@ if __name__ == "__main__":
# --- Control Arguments ---
cli_add_verbose(parser)
cli_add_version(parser, __name__)
args = parser.parse_args()
main(args)
+53
View File
@@ -0,0 +1,53 @@
#!/usr/bin/env python3
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# Tool to show the safetensors metadata
# python tool/show_metadata.py model.safetensors
import argparse
import json
from pprint import pprint
from seconohe.logger import logger_set_standalone
# Local imports
import bootstrap # noqa: F401
from src.nodes import main_logger
from src.nodes.utils.misc import cli_add_verbose, json_object_hook, cli_add_version
from src.nodes.inference.get_model import get_metadata
from src.nodes.inference.demucs_log_helper import DemucsModelInfo
def show_metadata(args):
d = get_metadata(args.input_file)
main_logger.info(f"Metadata information for `{args.input_file}`")
if 'model_t' not in d:
main_logger.warning("Missing model_t key, this isn't an AudioSeparation file")
expanded = {k: json.loads(v, object_hook=json_object_hook) if v[0] in '{[' else v for k, v in d.items()}
pprint(expanded)
sigs = expanded.get('signatures')
if sigs:
single = len(sigs) == 1
print("Demucs model")
if not single:
print(f"Composed by {len(sigs)} submodels")
weights = expanded.get('weights')
# Demucs model
for n, s in enumerate(sigs):
d = expanded[s]
num = -1 if single else n
DemucsModelInfo(num, d['class_name'], d['kwargs'], print, weights[n] if weights else None, extra=args.verbose,
sig=s)
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Shows the safetensors metadata",
formatter_class=argparse.RawTextHelpFormatter)
parser.add_argument('input_file', type=str, help="Path to the input safetensors model file.")
cli_add_verbose(parser)
cli_add_version(parser, __name__)
args = parser.parse_args()
logger_set_standalone(main_logger, args)
show_metadata(args)
Regular → Executable
+6 -3
View File
@@ -1,3 +1,4 @@
#!/usr/bin/env python3
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
@@ -10,11 +11,12 @@ import numpy as np
import onnx
from onnx import shape_inference # Import the shape inference module
import torch
from seconohe.logger import logger_set_standalone
import sys
# Local imports
import bootstrap # noqa: F401
from source.utils.logger import main_logger, logger_set_standalone
from source.utils.misc import cli_add_verbose
from src.nodes import main_logger
from src.nodes.utils.misc import cli_add_verbose, cli_add_version
def print_onnx_nodes_and_weights(onnx_model_path):
@@ -202,9 +204,10 @@ if __name__ == "__main__":
"Incompatible with -S")
parser.add_argument('-S', '--no_show', action='store_false', help="Don't print the ONNX structure.")
cli_add_verbose(parser)
cli_add_version(parser, __name__)
args = parser.parse_args()
logger_set_standalone(args)
logger_set_standalone(main_logger, args)
if args.run and not args.no_show:
main_logger.error("-r can't be used when -S is specified")
sys.exit(1)
Regular → Executable
+4 -3
View File
@@ -8,10 +8,11 @@
# python tool/uvr_hash.py model.onnx
import argparse
import os
from seconohe.logger import logger_set_standalone
# Local imports
import bootstrap # noqa: F401
from source.utils.logger import main_logger, logger_set_standalone
from source.db.hash import get_hash
from src.nodes import main_logger
from src.nodes.db.hash import get_hash
def main():
@@ -36,7 +37,7 @@ def main():
args = parser.parse_args()
args.verbose = 0
logger_set_standalone(args)
logger_set_standalone(main_logger, args)
# Process each file provided on the command line
for filepath in args.files: